diff --git a/Makefile b/Makefile index 6c8f3d5b042f..c9202e43a541 100644 --- a/Makefile +++ b/Makefile @@ -450,6 +450,7 @@ SOURCE_FILES = \ AlignLoads.cpp \ AllocationBoundsInference.cpp \ ApplySplit.cpp \ + Approximation.cpp \ Argument.cpp \ AssociativeOpsTable.cpp \ Associativity.cpp \ @@ -656,6 +657,7 @@ HEADER_FILES = \ AlignLoads.h \ AllocationBoundsInference.h \ ApplySplit.h \ + Approximation.h \ Argument.h \ AssociativeOpsTable.h \ Associativity.h \ @@ -2422,6 +2424,7 @@ install: $(LIB_DIR)/libHalide.a $(BIN_DIR)/libHalide.$(SHARED_EXT) $(INCLUDE_DIR cp $(ROOT_DIR)/tools/RunGenMain.cpp $(PREFIX)/share/halide/tools cp $(ROOT_DIR)/tools/halide_image.h $(PREFIX)/share/halide/tools cp $(ROOT_DIR)/tools/halide_image_io.h $(PREFIX)/share/halide/tools + cp $(ROOT_DIR)/tools/halide_approximation_testing.h $(PREFIX)/share/halide/tools cp $(ROOT_DIR)/tools/halide_image_info.h $(PREFIX)/share/halide/tools cp $(ROOT_DIR)/tools/halide_malloc_trace.h $(PREFIX)/share/halide/tools cp $(ROOT_DIR)/tools/halide_thread_pool.h $(PREFIX)/share/halide/tools diff --git a/apps/CMakeLists.txt b/apps/CMakeLists.txt index afc90b173081..70976a979c6f 100644 --- a/apps/CMakeLists.txt +++ b/apps/CMakeLists.txt @@ -53,6 +53,7 @@ add_app(cuda_mat_mul) add_app(depthwise_separable_conv) add_app(fft) add_app(gaussian_blur) +add_app(ggml) add_app(hannk) add_app(harris) # add_app(HelloAndroid) # don't build HelloAndroid here because it is driven by gradle diff --git a/apps/ggml/CMakeLists.txt b/apps/ggml/CMakeLists.txt new file mode 100644 index 000000000000..a06682271249 --- /dev/null +++ b/apps/ggml/CMakeLists.txt @@ -0,0 +1,81 @@ +cmake_minimum_required(VERSION 3.28) +project(ggml) + +enable_testing() + +set(CMAKE_CXX_STANDARD 17) +set(CMAKE_CXX_STANDARD_REQUIRED YES) +set(CMAKE_CXX_EXTENSIONS NO) + +# GGML is expected to be provided externally -- in this tree, via the +# apps/vcpkg/ports/ggml overlay port (see its portfile.cmake). GGML ships its +# own CMake package config (installed to /share/ggml/ by vcpkg's +# vcpkg_cmake_config_fixup), so no Find module is needed here. +find_package(ggml CONFIG REQUIRED) + +# Halide is already found once by the parent apps/CMakeLists.txt when this app +# is built as part of the full apps/ tree (a no-op re-find in that case); this +# call makes apps/ggml also independently configurable/buildable on its own, +# matching every other app's convention (see e.g. apps/blur/CMakeLists.txt). +find_package(Halide REQUIRED) + +# halide/ contains a from-scratch Halide reimplementation of GGML's Q4_0 +# quantize/dequantize kernels (ggml_quants_halide), benchmarked against +# GGML's own reference by providers/halide_provider.cpp below. +add_subdirectory(halide) + +add_executable( + kernel-bench + src/main.cpp + src/report.cpp + src/bench_quantize.cpp + src/bench_dequantize.cpp + src/bench_vecdot.cpp + src/bench_repack.cpp + providers/ggml_provider.cpp + providers/halide_provider.cpp +) + +target_include_directories(kernel-bench PRIVATE include providers src halide) + +target_link_libraries(kernel-bench PRIVATE ggml::ggml ggml_quants_halide) + +# GGML_VERSION isn't exposed through any runtime API, but ggml-config.cmake +# sets it as a plain CMake variable (baked in from the exporting build's own +# GGML_VERSION), so this is exactly the version of the library we just linked +# against -- reliable without touching GGML internals. +if (DEFINED GGML_VERSION) + target_compile_definitions(kernel-bench PRIVATE KERNEL_BENCH_GGML_VERSION="${GGML_VERSION}") +endif () + +# ggml-config.cmake.in only creates ggml:: import targets +# (including ggml::ggml-cpu) when GGML was built with GGML_BACKEND_DL=OFF +# (the default) -- see its `if (NOT GGML_BACKEND_DL)` guard. In DL mode the +# CPU backend is a runtime-loaded module with no link-time target at all, and +# the private ABI symbols this tool depends on (see +# providers/ggml_internal_abi.h) are then unreachable through the CMake +# package. Fall back to locating the library file directly by name; fail +# loudly if that isn't possible either, rather than producing a mysterious +# link error. +if (NOT TARGET ggml::ggml-cpu) + find_library( + GGML_CPU_DL_LIB + NAMES ggml-cpu + HINTS "${ggml_LIB_DIR}" "${ggml_LIB_DIR}/ggml" + PATH_SUFFIXES lib lib/ggml bin + ) + if (GGML_CPU_DL_LIB) + message( + STATUS "kernel-bench: linking ggml-cpu directly (GGML_BACKEND_DL build): ${GGML_CPU_DL_LIB}" + ) + target_link_libraries(kernel-bench PRIVATE "${GGML_CPU_DL_LIB}") + else () + message( + FATAL_ERROR "kernel-bench: could not find ggml::ggml-cpu or a standalone ggml-cpu library. " + "This GGML install appears to have been built with -DGGML_BACKEND_DL=ON, which " + "loads the CPU backend as a runtime module instead of linking it -- kernel-bench " + "needs to link directly against its internal symbols. Rebuild GGML with " + "-DGGML_BACKEND_DL=OFF (the default) and reinstall." + ) + endif () +endif () diff --git a/apps/ggml/PERF_NOTES.md b/apps/ggml/PERF_NOTES.md new file mode 100644 index 000000000000..9a57ba42047e --- /dev/null +++ b/apps/ggml/PERF_NOTES.md @@ -0,0 +1,394 @@ +# vec_dot performance notes + +Working notes for bringing `apps/ggml`'s vec_dot kernels up to `ggml-cpu` speed +on ARM (measured on an M3 Max). Branch `alexreinking/ggml-on-qk`, worktree +`~/dev/Halide/ggml-on-qk`. + +## Where things stand + +Measured at n=4096, best of many runs (see "Measuring" below): + +| type | ggml-cpu | halide | ratio | before this work | +| ---- | -------- | -------- | ----- | ---------------- | +| q4_0 | 87.8 ns | 90.7 ns | 0.97x | 0.97x | +| q4_1 | 106.7 ns | 110.2 ns | 0.97x | 0.80x | +| q5_0 | 122.2 ns | 130.2 ns | 0.94x | 0.51x | +| q5_1 | 143.5 ns | 153.7 ns | 0.93x | 0.50x | +| q8_0 | 69.1 ns | 70.4 ns | 0.98x | 0.95x | + +q5_K (still on the float path) also improved 8615 -> 6730 ns from the same +qh-expansion table (it shares the 1-bit `PlanarBitPack` decode). + +Everything else in the table is still on the unscheduled float path +(0.01x-0.12x) and untouched. All 28 roundtrip tests pass; `kernel-bench --all` +reports no mismatches; odd block-count tails are correct (verified at n = 32, +96, 160, 224, 1056). + +Commits, oldest first: + +- `c66ab336d` apps/ggml: bring q4_0/q8_0 vec_dot up to ggml-cpu speed +- `c8ec1934d` Add `Stage::distribute()`, and one accumulator per term in + `hoist_invariants()` +- `6c72fb4c1` apps/ggml: take the affine vec_dots (q4_1, q5_1) to SDOT +- `a6e9b0962` `hoist_invariants()`: return one Func per accumulator, not a Tuple + +## Dev aids + +The `getenv("GGML_PER_BLOCK_PROBE")` branch (original "variant A" via +`sdot_partial`) is kept, now reachable only for the *symmetric* SDOT formats +(q4_0/q8_0/q5_0) -- a dev aid to measure per-block vs the lane-split default. +The affine formats reach `sever_sum` first and never see it. + +Note the generator reads the env var at *generator run time*, so changing it +does not invalidate ninja's outputs. Force regeneration: + +```sh +rm -f build/apps/ggml/halide/q4_0_vec_dot.o build/apps/ggml/halide/libq4_0_vec_dot.a +GGML_PER_BLOCK_PROBE=1 cmake --build build/apps/ggml -j --target q4_0_vec_dot +cmake --build build/apps/ggml -j +``` + +## Building + +libHalide is consumed from `install/macOS`, so an app change needs only the app +build, but a Halide change needs build + install first: + +```sh +cmake --build build/macOS -j --target Halide +cmake --install build/macOS --prefix install/macOS +cmake --build build/apps/ggml -j +``` + +`build/macOS` is configured with `WITH_TESTS=OFF`, so `correctness_*` targets do +not exist. Compile a test directly instead: + +```sh +c++ -O1 -std=c++17 -DHALIDE_KEEP_MACROS -DHALIDE_WITH_EXCEPTIONS \ + -I install/macOS/include -I test/common -I tools \ + test/correctness/rfactor.cpp \ + -L install/macOS/lib -lHalide -Wl,-rpath,$PWD/install/macOS/lib -o /tmp/rfactor && /tmp/rfactor +``` + +`HALIDE_KEEP_MACROS` is required (`internal_assert` is `#undef`'d at the end of +the installed `Halide.h`); `HALIDE_WITH_EXCEPTIONS` is required or the +exception-guarded tests silently do not compile. To find which sub-test fails, +shard it: `TEST_TOTAL_SHARDS=40 TEST_SHARD_INDEX=N ./rfactor`. + +## Measuring + +macOS moves the process between P- and E-cores between runs, so a single run's +absolute numbers are not comparable -- swings of 30%+ are normal. Take the best +of several runs and always read the halide/ggml-cpu *ratio* within a run. + +`src/bench_vecdot.cpp` has two dev aids added this session: +`KERNEL_BENCH_FILTER=q4_0,q8_0` (substring match on type name) and +`KERNEL_BENCH_N=` (vector length; default 4096). The n sweep is what +separates per-call overhead from per-block cost -- fit the slope. + +Repeat-and-take-best wrapper: + +```sh +#!/bin/zsh +B=~/dev/Halide/ggml-on-qk/build/apps/ggml +N=${N:-5} +FILTER=${FILTER:-q4_0,q4_1,q8_0} +TMP=$(mktemp -d) +for i in $(seq $N); do KERNEL_BENCH_FILTER=$FILTER $B/kernel-bench --vecdot --csv $TMP/r$i.csv >/dev/null; done +cat $TMP/*.csv | awk -F, ' + $3!="role" && $1=="vec_dot" { k=$2 SUBSEP $4; if (!(k in best) || $5+0 < best[k]) best[k]=$5+0; ok[k]=$9; types[$2]=1 } + END { n=0; for (t in types) st[++n]=t + for(a=1;ast[b]){tmp=st[a];st[a]=st[b];st[b]=tmp} + for (i=1;i<=n;i++) { t=st[i]; c=best[t,"ggml-cpu"]; h=best[t,"halide"] + printf "%-8s %9.1f ns %10.1f ns %7.2fx %s\n", t, c, h, (h>0?c/h:0), (ok[t,"halide"]==1?"yes":"NO") } }' +rm -rf $TMP +``` + +To read generated code, re-run the generator by hand with extra outputs. Grab +the exact command from `build.ninja` (`grep 'COMMAND = .*-n q4_1_vec_dot '`) and +swap `-e c_header,object` for `-e stmt,assembly`. Add +`-no_asserts-no_bounds_query` to the target to see what actually ships. + +Diagnose SDOT vs fallback by grepping the `.s` for `sdot.4s`; grep the `.stmt` +for `vector_reduce_add(int32x..(widening_mul(int8x.., int8x..)))`. + +## What the q4_0/q8_0 speedup actually was + +Four independent things, roughly equal in size: + +1. **Lanes from `r.x`, not `r.y`.** `rfactor({{rxo, lane}, {r.y, u}})` keeps the + sdot's four Int(32) lanes alive into the float accumulator, so no block pays + a horizontal reduce. Lanes must come from `r.x`: blocks are interleaved + `{scale, codes}` records, so a lane per block gathers both the codes and the + scales. +2. **Chained sdot.** Cut `r.x` into chunks of 16 run serially, so both sdots + accumulate into the *same* register. Reducing straight to 4 lanes makes + `CodeGen_ARM` lower the wide reduce as two independent sdots plus an `addp` + (`codegen_dot_product_vector_reduce` only matches factor 4 and recurses). +3. **Interleave 4 blocks into independent accumulators.** Widening the vector + does not help -- every lane of one accumulator advances on every block, so + only interleaving blocks shortens the multiply-add chain. Un-interleaved, the + kernel is latency-bound at ~4 cycles/block. +4. **Per-call overhead.** vec_dot is called once per output element of a matvec, + so nothing amortizes. Three `Halide::Runtime::Buffer` constructions cost ~13 + ns flat; the assert/bounds-query prologue another ~10 ns. Fixed by filling a + `halide_buffer_t` in place (`StackBuffer` in `ggml_quants.cpp`) and building + these libraries with `FEATURES no_asserts no_bounds_query`. + +Two traps found along the way: + +- **A predicated tail is not a local cost.** Splitting the block RVar with + `GuardWithIf` makes the per-block sdot a dynamic-extent allocation that Halide + has to `bzero` and accumulate *through memory*, roughly doubling the cost of + every block. Fixed by giving the main reduction an exactly divisible extent + (`(nb / unroll_blocks) * unroll_blocks`) and sweeping the remainder in a + second update at the default schedule. `unroll_blocks` must be a power of two + -- 3 and 6 measured 40% worse because the simplifier cannot discharge the + tail. +- **`specialize()` inherits the schedule as of the call**, so scheduling + directives applied *after* `specialize()` do not reach the specialized branch. + It silently dropped the vectorize/unroll and made things 4x slower. + +### q4_0/q8_0 core-composition cleanup + +The tuned schedule above is now fed by faithful Approximation compositions (core +units plus the ggml units in `quant_components.h`) rather than ggml's legacy +`StructBlockLayout` and code-pack wrappers: + +- q4_0 `{Float16 d; UInt8 qs[16]}` is (in `Compose` encode order) + `BlockReshape -> SymmetricAffineQuantize -> Parallel{codes: nibble_offset(8) -> PlanarFieldPack, scale: fp16_storage()} -> StructLayout`. + `nibble_offset(8)` (an inline `Pointwise`) is the representation policy that + maps signed codes `[-8, 7]` to stored nibbles `[0, 15]`; planar packing + remains an exact, policy-free bit layout. +- q8_0 `{Float16 d; Int8 qs[32]}` is + `BlockReshape -> SymmetricAffineQuantize -> Parallel{scale: fp16_storage()} -> StructLayout`. + Its faithful signed array means the old UInt8 `BytePack` reinterpretation is + unnecessary. + +Both formats share one traced weight/activation decode graph between the main +and remainder updates. The shared stages are eagerly inlined only into the tail +before the four-block main update is transformed. The main assembly remains +eight SDOTs per iteration with paired 128-bit code loads, four persistent vector +accumulators, an explicitly unrolled epilogue, and no accumulator spill. The +scalar tail is still intentionally proportional to its one-to-three-block +remainder; include those cases in the planned named scaling/odd-tail +`kernel-bench` mode. + +A struct-typed q8_0 activation input was also tested. It was correct, but it +broadened load-shape changes across q4_0 and q5_0 without removing any target +weight/codec representation debt, so the shared activation ABI remains on its +stable byte path. The standalone q8_0 codecs and q8_0 weight operand use the +faithful struct composition. Q1_0 and struct-aware reblocking are the remaining +symmetric compatibility work. + +Ten final paired n=4096 runs measured q4_0 at 95.088 ns GGML / 97.000 ns Halide +(0.9803x), versus a worktree baseline of 92.585 / 95.784 ns (0.9666x). Q8_0 +measured 72.973 / 74.951 ns (0.9736x), versus 70.618 / 74.870 ns (0.9432x). +Absolute times moved with core placement; the paired ratios improved, and Halide +time changed by +1.27% and +0.11%, respectively. No compiler change was +required: the simplifier folded q4's offset/planar stages into the expected +mask, shifts, and vector add, while q8's signed field lowered to direct vector +loads. + +## The q4_1 story so far + +q4_1 is affine: the weight decodes to `d*code + m`, so the per-block product +`(d*code + m) * (d_act*act)` has no single scale to hoist and it was left on the +default (fully scalar, unscheduled) float reduction at 3843 ns. + +`Stage::distribute()` multiplies the product out to +`d*d_act * sum(code*act) + m*d_act * sum(act)`, and `hoist_invariants()` gives +each term its own accumulator. Both bodies are integer, so both reach SDOT -- +the first as the ordinary dot, the second as a dot with a vector of ones (ARM +already matches `i32(int8x)`). That is ggml's own decomposition, and it got q4_1 +to 133.8 ns / 0.80x. + +**Design decisions worth not re-litigating:** + +- Multiplying out is `distribute()`, a separate schedule directive, *not* + something `hoist_invariants()` decides. A first attempt did it unconditionally + and broke `hoist_invariants test (predicated RDom)`: `require(...) * (r + 1)` + distributes into two accumulators when one was optimal. The predicate that + would have rescued it ("ignore constants, don't traverse call arguments...") + is exactly the kind of heuristic that needs tuning forever. Whether to + multiply out depends on what the terms turn out to contain, which is a cost + question, so it belongs in the schedule. +- `hoist_invariants()` returns **one single-valued Func per accumulator**, not a + Tuple. The Tuple version worked and measured identically, but it blocked + `sever` (which cannot sever one value of a Tuple) and forced `change_type` to + grow multi-output support. Fusing the terms' loop nests is `compute_with`'s + job. + +## The stored block sum -- DONE (q4_1 0.97x) + +ggml does *not* recompute `sum(act)` at vec_dot time -- it reads the `s` field +that `block_q8_1` stores at quantize time +(`{ggml_half d; ggml_half s; int8_t qs[32]}`, 36 bytes). Our Q8_1 codec already +computes that field (`AppendSums{block_size, SumMode::ScaledFloat}` in +`make_symmetric_byte_sum_block_scheme`). We now sever the offset term's +accumulator straight to it via `sever`, which needs no new Halide directive: its +contract already *is* this claim -- sever a Func's computation, replace calls +with an ImageParam read, discard the recomputing reduction. + +**Result: variant A + sever measured 110.2 ns / 0.97x** (q5_1: 208 ns / 0.64x). +This *matched* the lane-split+sever projection, so the second route below +(teaching `distribute()` to split into separate update definitions for per-term +rfactors) is **not needed** -- skip it. All correct: 28 roundtrips, +`kernel-bench --all` clean, odd tails at n = 32/96/160/224/1056. + +The recipe, as implemented in `vec_dot_generator_base.h`'s `sever_sum` branch +(reached when `distribute_terms && act_has_block_sums`, i.e. affine x Q8_1): + +1. `acc_dot = Acc.update().rfactor({{r.y, u}})` -- whole-block partials (variant + A). Inline **only the weight's** decode chain (replacement + inlinable + handles, multi-pass), leaving the activation decode `Act` + (`act_r.replacement`) whole. `distribute()`, then `hoist_invariants()`. The + offset term's accumulator body is then `Act(r.x, u)`, so the accumulator *is* + `sum_k decode_act(k, blk)` = the stored `s`. +2. `parts[1].change_type(Float(16))` -- makes the severed Func's type match the + data. Faithful: the encoder rounds `s` to fp16 too (the ~6e-06 rel err is + exactly that rounding, same as ggml's). +3. `Pipeline({Acc}).sever({s16}, {s_blocks})` -- second `sever` on the pipeline + (the first, at configure top, severs the encode halves to x_blocks/y_blocks). + `s_blocks` is the third Input. +4. Product term: inline `Act`'s **full** chain into `parts[0]` (one Act + eager_inline is not enough -- the chain has intermediate levels, and a single + inline leaves the second hoist with no visible d_act factor and it errors), + re-`hoist_invariants()` to pull d_act out, `change_type(Int(32))` -- the + survivor reaches SDOT. Verified: 8 `sdot` in the `.s`; writeback is + `d_w*(d_act*int32_dot) + m_w*s_blocks[blk]`, ggml's exact decomposition, with + no `sum(act)` reduction left anywhere in the stmt. + +Plumbing (all landed, see Uncommitted): + +- `s_blocks`: 1-D `Float(16)` Input, `dim(0).set_stride(act_bytes/2)` (= 18 for + Q8_1) -- pinning it makes the read an immediate offset, *and* is required: + left dynamic, Halide's default constrains the innermost stride to 1 and the + bound-check fails against the strided view. +- ABI: `StackBuffer::blocks_field_f16(base, nb, byte_offset=2, block_bytes=36)` + -- a zero-copy fp16 view of the `s` slot, stride `block_bytes/2`. +- Format knowledge lives with the codec: `SchemeAndBytes::has_block_sums` set by + `make_symmetric_byte_sum_block_scheme`, carried to + `VecDotSpec::act_has_block_sums` (guarded `&& a_nat == wbs`, so a Reblock'd + activation stays off it). + +**The structural conflict (why variant A, not lane-split).** The two terms want +different rfactors: `sum(code*act)` wants `{rxo->lane, r.y->u}` (the lane split +is the q4_0/q8_0 win); `sum(act)` must be the *whole-block* sum to equal `s`, so +`{r.y->u}` only. One update rfactors one way. Variant A gives both `{r.y->u}`; +severing then deletes the offset accumulator entirely, so its per-lane-partial +problem never arises -- and the surviving product dot, alone in its block loop, +schedules close enough to the lane-split base that the ~8 ns gap projected +between the routes did not materialize. `compute_with` (verified to fuse a +lane-split with a block-only reduction using `AlignStart` + matching split-var +names) is therefore unnecessary here; keep it in mind for formats that keep two +live accumulators. + +## The q5_x 5-bit high bit -- DONE (q5_0 0.94x, q5_1 0.93x) + +q5_0/q5_1 reach SDOT, but every code carries a per-element high bit unpacked +from the `qh` field, and that reconstruction -- not the dot, not (for q5_1) the +offset -- is the whole gap vs q4_x. + +**What the reconstruction was.** `PlanarBitPack::decode`'s 1-bit case emitted +`(qh[kk/8] >> (kk%8)) & 1`: a per-lane variable shift plus a `transpose_vector` +to broadcast the two `qh` bytes across the 16 sdot lanes +(`dup.8b + dup.4h + uzp2` per sdot, ~24 NEON ops/block). ggml avoids this with a +1 KB `table_b2b` memory LUT (byte -> 8 expanded bytes, one contiguous load per +`qh` byte). + +**What was done.** Mirrored the LUT: compile-time b2b tables embedded in the +binary. The q5 tables contain the final high-bit contribution (`-16/0` for q5_0, +`0/16` for q5_1), so the lookup folds the shift and q5_0 zero-point into the +load rather than reconstructing a raw bit and applying them afterward. Several +other details are necessary: + +1. The table read is only a *contiguous* 8-byte load + (`b0[ramp(qh_byte*8, 1, 8)]`, matching ggml) when the `qh` byte is a + **scalar** and the 8 bit positions are the vector lanes. Inlined into the + sdot it is the opposite (the byte varies per lane -> a 16-wide per-lane + gather, `ld1` per lane, measured 0.12x). So the reconstructed codes are + **materialized** per block (`combine_bits_code` compute_at the block loop, + `kk` split `(byte, pos)`: pos vectorizes the load, byte unrolls to a scalar + index). To keep the codes in the vector register file rather than round-trip + a stack buffer, the materialization is `store_in(MemoryType::Register)` with + `kk` split further into 16-code units (one sdot chunk = two `qh` bytes x 8 + positions) so the store width matches the sdot's read width -- otherwise the + 8-wide table-load store vs 16-wide sdot load mismatch keeps it in memory. +2. Materializing needs the codes leaf kept out of `sdot_partial()`'s deep inline + (`can_be_inlined()` ignores compute level, so a `compute_root`/`compute_at` + schedule alone does not stop `eager_inline`). `sdot_partial()` now takes + resolved `Func` identities, not generated-name strings; q5_0 obtains the + identity from its `AdditiveRadixSplit` stage key. +3. q5_0 uses its faithful packed type `{Float16 d; UInt8 qh[4]; UInt8 qs[16]}`. + `LittleEndianScalarPack` decodes qh with `concat_bits()`; + `LowerStructTypes` and LLVM consequently issue one unaligned `ldr/ldur w`, + matching ggml, rather than four `ldrb`s. q5_1 retains the older scalar + `UInt(32)` struct field as explicit transitional debt. +4. q5_0's odd-block **tail** now shares the main Approximation graph. Before the + main update materializes the keyed reconstructed-code and qh-word Funcs, the + tail update alone eagerly inlines that decode chain. This removes the second + q5_0 weight `approximate_by` graph while keeping the tail inline and the main + loop register-resident. q5_1 retains its separate legacy tail chain. + +The transpose is gone (verified: no `uzp2`/`dup.8b`); reconstruction is a +handful of ops + 4 contiguous LUT loads/block. Bonus: the same decode change +sped up q5_K on the float path (8615 -> 6730 ns). + +**The schedule changes still matter after the compiler fix.** The compiler +change is the dominant improvement, but it does not make the schedule work +redundant: + +- Removing the explicit qh `compute_at(...).store_in(Register)` schedule made + q5_0 regress from about 122 to 129 ns and q5_1 from about 152 to 162 ns, even + though both generated versions contained a word load. Keep it: it changes + placement/reuse around the four bit extracts, not merely the load width. +- Interleaving two q5 blocks gives better latency hiding/register pressure than + the four-block setting used by q4/q8. Applying two globally regressed q4_0 and + q4_1, so this is a per-format `VecDotSpec` choice. +- Explicitly unrolling the fixed-size final rfactor reductions removes a stack + spill and one-iteration epilogue loops (roughly another 1 ns here). + +**q5_1 needs a different reduction shape.** The generic affine path horizontally +reduced each block's Int32 dot before accumulating floats. The q5_1-specific +schedule keeps four float dot lanes live across the entire block loop, maintains +two scalar `m*s` accumulators alongside them, and fuses their two-block loops. +The generated steady state now matches the important shape of ggml's kernel: two +SDOTs per block, persistent vector FMAs, and scalar offset FMAs, with one +horizontal reduction at the end. + +After the core-composition migration, ten paired n=4096 runs had q5_0 medians of +122.147 ns GGML / 130.645 ns Halide and a median paired ratio of 0.9349x. The +committed baseline was 122.172 / 130.259 ns and 0.9379x, so the refactor is +performance-neutral (+0.30% Halide time). q5_1 remained within noise at 143.728 +/ 153.956 ns and 0.9336x. Use repeated paired runs rather than treating either +sample as a fixed score. + +**Core composition now mirrors the representation.** q5_0 is +`BlockReshape -> SymmetricAffineQuantize -> Parallel{codes: AdditiveRadixSplit} -> Parallel{low: PlanarFieldPack, high: BinaryAlphabetPack -> LittleEndianScalarPack, scale: fp16_storage()} -> StructLayout` +(in `Compose` encode order). Opaque Approximation stage keys carry the +reconstructed-code and qh-word identities into `VecDotSpec`; generated Func +names no longer control q5_0's schedule. `BlockReshape` lives in the core +Approximation library; symmetric block quantization, `AdditiveRadixSplit` and +`BinaryAlphabetPack` live in ggml's `quant_components.h`. + +Reusable performance experiments should become discoverable `kernel-bench` +modes. In particular, promote the odd-tail n sweep (32, 96, 160, 224, 1056) +instead of preserving it only as a shell loop. + +**Build-flag bug.** q5_0/q5_1's `add_halide_library` were missing +`FEATURES no_asserts no_bounds_query` (every other tuned vec_dot has it), so +they paid the assert/bounds-query prologue -- ~11 ns of pure startup on a ~170 +ns call. Adding it was worth 0.63 -> 0.67x on its own; the register store +another 0.67 -> 0.69x (q5_1 to 0.78x). Check this first on any new kernel. + +Group-level codes compute (all blocks first) measured worse than the final +per-block materialization. Likewise, removing qh materialization because the new +compiler lowering already produced `ldr w` was a measured regression; the two +optimizations address different parts of the generated loop. + +## Not started + +Everything else (k-quants, IQ family, tq\*) is still on the default unscheduled +float reduction and would benefit from the same treatment as q4_1 -- the +k-quants' two-level scales are also sums of scaled sub-reductions, which is what +`distribute()` was built for. diff --git a/apps/ggml/Q4_0_Q8_0_APPROXIMATION_PLAN.md b/apps/ggml/Q4_0_Q8_0_APPROXIMATION_PLAN.md new file mode 100644 index 000000000000..2edb1b85e2d8 --- /dev/null +++ b/apps/ggml/Q4_0_Q8_0_APPROXIMATION_PLAN.md @@ -0,0 +1,140 @@ +# Clean q4_0 and q8_0 Approximation Implementation + +## Current status + +- Source baseline commit: `ca685e94e91c33d924f10c238ffa4ae6dab183a4` +- Working baseline: completed q5_0 core-composition refactor in this worktree +- Status: complete; all acceptance gates passed +- [x] Phase 0: inspect the existing schemes and create this progress document +- [x] Phase 1: collect ten paired baseline runs at n=4096 +- [x] Phase 2: compose q4_0 and q8_0 from reusable core components +- [x] Phase 3: add focused correctness coverage +- [x] Phase 4: preserve identity-based scheduling and tuned tails +- [x] Phase 5: correctness, odd-tail, generated-code, and full-suite validation +- [x] Phase 6: paired performance validation and durable documentation + +## Goal and constraints + +Apply the q5_0 cleanup methodology to q4_0 and q8_0. Their schemes should use +faithful packed struct types and reusable public Halide Approximations; the +generic vec-dot generator should retain only the base reduction, +`approximate_by`, `sever`, and scheduling. Preserve bit-exact quantization, +dequantization/vec-dot correctness, and the tuned four-block SDOT shape. Keep +median paired performance within 5% of this worktree baseline and at least 0.90x +GGML. + +Unrelated legacy formats remain compatibility debt. In particular, q1_0 keeps +the legacy symmetric layout until it is migrated deliberately, and q4_1/q5_1 +remain on their affine/legacy paths. + +## Target compositions + +`Compose` lists stages in encode order (innermost first). + +q4_0 faithful type: `{d: Float16, qs: UInt8[16]}`. + +1. `BlockReshape{32}` with the requested row/block-indexed layout. +2. `SymmetricAffineQuantize`, qmax 8, extreme-signed scale selection, and + truncate-half-up-with-offset rounding. +3. `Parallel` on `codes`: an inline `Pointwise` additive offset + (`nibble_offset`) mapping signed codes to on-disk `[0, 15]` values, then + `PlanarFieldPack{4, 16}` for stored nibbles; and on `scale`: an inline + `Pointwise` storage cast (`fp16_storage`) from Float32 to Float16. +4. `StructLayout`, logical `{qs, d}` to physical fields. + +q8_0 faithful type: `{d: Float16, qs: Int8[32]}`. + +1. `BlockReshape{32}` with the requested row/block-indexed layout. +2. `SymmetricAffineQuantize`, qmax 127, absolute-max scale selection, and + nearest rounding. +3. `Parallel` on `scale`: `fp16_storage` (Float32 to Float16). +4. `StructLayout`, logical `{qs, d}` to physical fields. + +The standalone q8_0 codecs and q8_0 weight path use the faithful core scheme. +The shared activation ABI remains byte-addressed: a struct-typed experiment +broadened generated-code changes across q4_0/q5_0 without removing reusable +representation logic from either target's weight/codec pipeline. Mismatched +consumers also require that byte path for the existing `Reblock` component. + +## Validation gates + +- q4_0 and q8_0 quantize outputs are bit-exact with GGML; dequantize and vec-dot + pass existing tolerances. +- Focused component tests cover the additive offset and both compositions. +- `kernel-bench --all` has no failures. +- Odd block counts pass at n=32, 96, 160, 224, and 1056. +- ARM main loops retain SDOT, four blocks in flight, wide contiguous code loads, + persistent accumulators, and fully unrolled fixed-size epilogues without + accumulator stack spills or one-iteration epilogue loops. +- Median paired performance is no more than 5% slower than baseline and remains + at least 0.90x GGML; q5_0 and affine shared-format checks show no accidental + regression. + +## Benchmark experiment policy + +Any useful or repeatable experimental setup must be promoted into `kernel-bench` +as a named mode rather than left as an ad hoc shell recipe. Size and scaling +sweeps used here should feed the same scaling/odd-tail mode already identified +by the q5_0 work. + +## Baseline results + +Ten paired filtered runs at `KERNEL_BENCH_N=4096`: + +| Format | GGML CPU | Halide | Paired GGML/Halide | +| ------ | ---------- | ---------- | ------------------ | +| q4_0 | 92.585 ns | 95.784 ns | 0.9666x | +| q4_1 | 104.726 ns | 118.332 ns | 0.8850x | +| q5_0 | 119.370 ns | 128.931 ns | 0.9258x | +| q5_1 | 135.748 ns | 153.946 ns | 0.8818x | +| q8_0 | 70.618 ns | 74.870 ns | 0.9432x | + +All candidate correctness flags were true. Raw CSV files are in +`/tmp/q48-baseline.4xCRUB` for this work session. + +## Experiment log + +| # | Change | Correctness | GGML / Halide timings | Generated-code observations | Decision | +| --- | ----------------------------------------------------------------------------------------------------------------------------------- | --------------------------------------------------------- | ----------------------------------------------------------------------------------------------------- | -------------------------------------------------------------------------------------------------------- | ------------------------------------------------ | +| 0 | Completed q5_0 worktree baseline | All filtered vec-dot checks passed | q4_0: 92.585 / 95.784 ns, 0.9666x; q8_0: 70.618 / 74.870 ns, 0.9432x | Existing four-block SDOT paths are the generated-code reference | Reference | +| 1 | Add an inline `Pointwise` additive offset (`nibble_offset`) and compose faithful q4_0/q8_0 structs from public components | Component composition and standalone q4_0/q8_0 tests pass | Ten-run probe: q4_0 94.177 / 96.975 ns, 0.9712x; q8_0 72.942 / 74.824 ns, 0.9748x | Core stages simplify to the existing signed-code SDOT inputs | Keep | +| 2 | Use faithful struct q8_0 for the shared activation operand | Correct, including q5_0 | Single sample moved absolute timings with core placement; paired ratios did not indicate a regression | Broadened input/load-shape changes across q4_0 and q5_0 | Revert; keep the established byte activation ABI | +| 3 | Share the traced weight and activation decode graphs between q4_0/q8_0 main and tail updates, eagerly inlining only the tail update | Standalone and odd-size tests pass | Ten-run probe: q4_0 96.039 / 99.147 ns, 0.9687x; q8_0 75.801 / 77.443 ns, 0.9788x | Removes duplicate tail Approximation graphs; main remains four-block SDOT and the remainder stays scalar | Keep | +| 4 | Full validation and ten final paired runs | `kernel-bench --all` clean; all paired flags true | q4_0: 95.088 / 97.000 ns, 0.9803x; q8_0: 72.973 / 74.951 ns, 0.9736x | Eight SDOTs/four blocks, paired 128-bit code loads, persistent accumulators, no accumulator spill | Final | + +## Final paired results + +Median of ten n=4096 paired runs: + +| Format | GGML CPU | Halide | Paired GGML/Halide | Halide vs baseline | +| ------ | ---------- | ---------- | ------------------ | ------------------ | +| q4_0 | 95.088 ns | 97.000 ns | 0.9803x | +1.27% | +| q4_1 | 106.761 ns | 117.951 ns | 0.9051x | -0.32% | +| q5_0 | 122.167 ns | 130.589 ns | 0.9355x | +1.29% | +| q5_1 | 143.560 ns | 153.727 ns | 0.9339x | -0.14% | +| q8_0 | 72.973 ns | 74.951 ns | 0.9736x | +0.11% | + +Negative deltas are improvements. Raw final CSV files are in +`/tmp/q48-final.5IktYc` for this work session. + +## Framework/compiler issues + +- No compiler change was needed. The simplifier folds q4_0's core + `nibble_offset` and `PlanarFieldPack` into the same mask/shift/vector-add + operations consumed by SDOT, and q8_0's signed struct array lowers to direct + 128-bit loads. +- A struct-typed q8_0 activation was correct but unnecessarily broadened load + shape changes across q4_0/q5_0. The stable shared activation ABI remains the + compatibility byte path; q8_0's codecs and weight path are fully + core-composed. +- Shared q4_0/q8_0 tails must be eagerly inlined into the tail update before the + main SDOT schedule is applied. The remainder is deliberately scalar and its + cost scales with one to three blocks; this reinforces the need for a named + size-sweep benchmark mode. + +## Final follow-up items + +- Add the reusable size/scaling experiments from this work as named + `kernel-bench` modes. +- Migrate q1_0 from `StructBlockLayout` and make `Reblock` struct-aware before + removing the symmetric compatibility layout and byte activation path. diff --git a/apps/ggml/Q5_0_APPROXIMATION_PLAN.md b/apps/ggml/Q5_0_APPROXIMATION_PLAN.md new file mode 100644 index 000000000000..cb85b9eb155e --- /dev/null +++ b/apps/ggml/Q5_0_APPROXIMATION_PLAN.md @@ -0,0 +1,152 @@ +# Clean q5_0 Approximation Implementation + +## Current status + +- Baseline commit: `ca685e94e91c33d924f10c238ffa4ae6dab183a4` +- Status: complete; all acceptance gates passed +- [x] Phase 0: confirm clean baseline and create this progress document +- [x] Phase 1: collect ten paired baseline runs at n=4096 +- [x] Phase 2: add stage-key tracing and reusable core Approximation components +- [x] Phase 3: add focused core correctness coverage +- [x] Phase 4: refactor q5_0 composition and schedule lookup +- [x] Phase 5: correctness, odd-tail, generated-code, and full-suite validation +- [x] Phase 6: paired performance validation and durable documentation + +## Goal and constraints + +Refactor q5_0 so its generator contains only the base reduction, +`approximate_by`, `sever`, and scheduling. Compose all q5_0 representation logic +from reusable core Halide Approximations. Preserve correctness, keep median +paired performance within 5% of the committed baseline, and remain at least +0.90x GGML. q5_1 is explicitly unchanged transitional debt. + +## Core Approximation APIs + +- Add an opaque, copyable `ApproximationStageKey` to every Approximation + instance. +- Extend traced invocation through `Compose`, `Parallel`, `TrustedInverse`, and + `Func::approximate_by` so `ApproximationResult` resolves encoded and decoded + outputs by `(StageKey, port)` while preserving flat `handles` and `encoded`. +- Add public standard components in a dedicated header included by `Halide.h`: + `StructLayout`, `LittleEndianScalarPack`, `PlanarFieldPack`, `BlockReshape`, + and an inline `Pointwise` storage cast. (`BinaryAlphabetPack`, + `AdditiveRadixSplit`, and the symmetric block quantizer/policies were later + moved into ggml's `quant_components.h`.) +- Leave ggml compatibility aliases/wrappers so unrelated formats do not migrate. +- Test nested combinators, repeated types with distinct keys, invalid lookups, + encode/decode lookup, and component correctness, including one- and + two-dimensional `StructLayout` records. + +## q5_0 target composition and schedule + +Faithful packed type: `{d: Float16, qh: UInt8[4], qs: UInt8[16]}`. + +In `Compose` encode order (innermost first): + +1. `BlockReshape{32, block_indexed}`. +2. Symmetric block quantization, qmax 16, extreme-signed scale selection, + truncate-half-up rounding. +3. `Parallel` on `codes`: `AdditiveRadixSplit{16, 16}` splits signed codes into + `low` and `high`. +4. `Parallel` on the resulting ports: `PlanarFieldPack{4, 16}` for `low` + nibbles; `BinaryAlphabetPack{32, UInt32, -16, 0}` then + `LittleEndianScalarPack` for `high` contributions; `fp16_storage` + (Float32 to Float16) for `scale`. +5. `StructLayout`, logical `{qs, qh, d}` to physical fields. + +Capture stage keys for reconstructed signed codes (`AdditiveRadixSplit`) and the +qh word (`LittleEndianScalarPack`), carry them through scheme metadata into +`VecDotSpec`, and resolve scheduling Funcs from `ApproximationResult`. Replace +the duplicate q5_0 tail Approximation with stage-scoped eager inlining into the +tail update. Preserve the measured two-block SDOT schedule and use Func +identity, not names, for `sdot_partial` exclusions. + +## Validation gates + +- q5_0 quantize is bit-exact with GGML; dequantize and vec_dot pass tolerances. +- `kernel-bench --all` has no failures. +- Odd block counts pass at n=32, 96, 160, 224, and 1056. +- ARM assembly contains SDOT, one qh word load per block, contiguous LUT loads, + no qh byte-load sequence in the main loop, and no accumulator stack spill or + one-iteration epilogue loop. +- Median paired q5_0 is no more than 5% slower than baseline and at least 0.90x + GGML; shared formats show no accidental regression. +- Any compiler optimization is general, separately tested, and + target-independent where possible. No new scheduling directive is planned. + +## Benchmark experiment policy + +Any experimental setup that proves useful or repeatable should be promoted into +`kernel-bench` as a named mode rather than left as an ad hoc shell recipe. This +includes scaling/size sweeps such as the q5_0 odd-block checks at n=32, 96, 160, +224, and 1056. One-off scripts are acceptable for initial exploration, but the +durable form should make the experiment discoverable, reproducible, and usable +for future regressions from the benchmark utility itself. + +## Baseline results + +Ten paired filtered runs, `KERNEL_BENCH_N=4096`, filter +`q4_0,q4_1,q5_0,q5_1,q8_0`: + +Median of ten per-run timings and median paired ratio: + +| Format | GGML CPU | Halide | Paired GGML/Halide | +| ------ | ---------- | ---------- | ------------------ | +| q4_0 | 94.437 ns | 96.893 ns | 0.9726x | +| q4_1 | 107.819 ns | 118.424 ns | 0.9020x | +| q5_0 | 122.172 ns | 130.259 ns | 0.9379x | +| q5_1 | 143.575 ns | 153.717 ns | 0.9339x | +| q8_0 | 73.202 ns | 74.829 ns | 0.9764x | + +All candidate correctness flags were true. Raw CSV files are in +`/tmp/q50-baseline.tX9dNF` for this work session. + +## Experiment log + +| # | Change | Correctness | GGML / Halide timings | Generated-code observations | Decision | +| --- | ----------------------------------------------------------------------------------------------------------------------------------------- | -------------------------------------------------------------------------------------------------------- | ----------------------------------------------------------------- | -------------------------------------------------------------------------------------------------------------------- | --------- | +| 0 | Clean committed baseline | All filtered vec_dot checks passed | q5_0: 122.172 / 130.259 ns, 0.9379x paired median | Existing q5_0 path is the generated-code reference | Reference | +| 1 | Add opaque stage keys and traces through Compose, Parallel, TrustedInverse, and approximate_by | Nested/repeated/directional/invalid lookup tests pass | Not performance-sensitive | Flat encoded/handle compatibility retained | Keep | +| 2 | Add public core components; move BlockReshape and symmetric quantization behind ggml compatibility alias/wrapper (since re-homed in ggml) | Focused component suite passes, including 1-D and 2-D StructLayout and inline fp16 storage-cast rounding | Not measured independently | `strict_float` makes fp16 storage rounding survive eager inlining | Keep | +| 3 | Replace q5_0 legacy split-code composition with the eight reusable stages | q5_0 quantize/dequantize/vec_dot pass | Focused q5_0 run: 117.7 / 126.2 ns in one sample | Faithful UInt8[4] qh lowers through concat_bits to a word load | Keep | +| 4 | Resolve codes/qh by stage key, stop sdot inlining by Func identity, and share/eager-inline the q5_0 tail decode graph | Odd n=32/96/160/224/1056 all pass | Included in final medians | Main loop has two blocks in flight, four SDOTs/pair, persistent vector accumulators; tail is scalar and fully inline | Keep | +| 5 | Full validation and ten final paired runs after the final storage-cast change | `kernel-bench --all` clean; focused tests and odd sizes pass; all paired flags true | q5_0: 122.147 / 130.645 ns, 0.9349x; baseline delta +0.30% Halide | One unaligned qh word load/block, contiguous LUT loads, no accumulator spill or one-iteration epilogue | Final | + +## Final paired results + +Median of ten n=4096 paired runs: + +| Format | GGML CPU | Halide | Paired GGML/Halide | Halide vs baseline | +| ------ | ---------- | ---------- | ------------------ | ------------------ | +| q4_0 | 94.187 ns | 96.832 ns | 0.9727x | -0.06% | +| q4_1 | 106.760 ns | 118.423 ns | 0.9015x | 0.00% | +| q5_0 | 122.147 ns | 130.645 ns | 0.9349x | +0.30% | +| q5_1 | 143.728 ns | 153.956 ns | 0.9336x | +0.16% | +| q8_0 | 72.774 ns | 74.687 ns | 0.9744x | -0.19% | + +Negative deltas are improvements. Raw final CSV files are in +`/tmp/q50-final-strict.vULMku` for this work session. + +## Framework/compiler issues + +- No compiler change was needed. Existing `LowerStructTypes` handling of + `concat_bits()` recovered one unaligned qh word load from the faithful + `UInt8[4]` field. +- The existing Halide build had `WITH_TESTS=OFF`; it was reconfigured with tests + enabled to build the two focused correctness targets. +- Installing only the development component updates headers and GenGen but not + the changed shared library; a full `cmake --install` was required before the + standalone ggml generator could link the new stage lookup methods. +- A plain fp32-to-fp16-to-fp32 cast chain can fuse away when fully inlined. The + fp16 storage cast (`fp16_storage`) uses `strict_float` on both conversions so + storage rounding is schedule-independent; its correctness test intentionally + leaves the stage inline. q5_0 also materializes the fp16 value in its packed + struct field. + +## Final follow-up items + +- Add a `kernel-bench` scaling/odd-tail mode that captures the useful n-sweep + performed during this work; follow the benchmark experiment policy above for + future reusable setups. +- Remove the legacy q5-specific components and specialized q5_1 reduction after + q5_1 is migrated to reusable components; this work intentionally leaves it. diff --git a/apps/ggml/README.md b/apps/ggml/README.md new file mode 100644 index 000000000000..1d5adc675b90 --- /dev/null +++ b/apps/ggml/README.md @@ -0,0 +1,66 @@ +# kernel-bench + +Benchmarks GGML's CPU quantize / dequantize / vec_dot / repack kernels against a +designated correctness reference, and is structured so that a from-scratch +implementation of any of those kernels can be dropped in and compared too. See +`providers/README.md` for how to add one. + +This is a **standalone** CMake project. It is not built as part of GGML itself +and consumes an already-built-and-installed GGML purely as an external +dependency via `find_package`. + +## Build + +```sh +# 1. Build and install GGML somewhere (skip if you already have an install). +# GGML_BACKEND_DL=OFF (the default) is required -- see the "private ABI" note below. +cmake -S /path/to/ggml -B /path/to/ggml/build -DCMAKE_BUILD_TYPE=Release +cmake --build /path/to/ggml/build -j +cmake --install /path/to/ggml/build --prefix /path/to/ggml/install + +# 2. Build kernel-bench against that install. +cmake -S . -B build -DCMAKE_PREFIX_PATH=/path/to/ggml/install +cmake --build build -j + +./build/kernel-bench --all +``` + +## What the report means + +Each row shows a `ggml_type` (or, for repack, a specific interleave layout like +`q4_0_4x4_q8_0`), GGML's designated **reference** implementation and its +timing/throughput, and one column per **candidate** implementation registered +for that kernel: its timing/throughput, speedup relative to the reference, and +whether its output matched the reference (quantize/repack packing is checked +byte-for-byte; dequantize/vec_dot/gemv/gemm results are checked within a +relative-error tolerance, since those involve floating point accumulation that +different implementations may order differently). + +A candidate flagged "identical to reference" has the exact same function address +as the reference -- this happens whenever the current CPU architecture has no +separate optimized kernel for that type (GGML's `src/ggml-cpu/arch-fallback.h` +collapses the two names onto one symbol in that case), so timing it separately +would only measure noise. + +The dequantize table currently shows only a reference column with no candidates: +GGML has exactly one dequantize implementation per type (arch-independent, in +`src/ggml-quants.c`), so there's nothing to compare it against yet -- this is +intentionally the first place to plug in a new provider (see +`providers/README.md`). + +## Why this needs a private ABI header + +GGML's public API (`ggml_get_type_traits` / `ggml_get_type_traits_cpu` in +`include/ggml.h` / `include/ggml-cpu.h`) exposes exactly one reference and one +CPU-dispatched implementation per type, which is sufficient for the quantize and +dequantize benchmarks without touching anything private. It does **not** expose +the always-available pure-C fallback for `vec_dot`, nor anything for the repack +`quantize_mat`/`gemv`/`gemm` kernels. Those are only reachable because +`ggml-cpu` is built without `-fvisibility=hidden`, so its internal (but +non-`static`) C symbols end up with default/exported linker visibility by +accident of the build configuration rather than by design. +`providers/ggml_internal_abi.h` redeclares exactly the symbols needed, copied +from GGML's uninstalled `src/ggml-cpu/quants.h` / `repack.h` as of the commit +this tool was written against. If a future GGML release renames or changes the +signature of one of these functions, that header (and +`providers/ggml_provider.cpp`) are the only places that need updating. diff --git a/apps/ggml/halide/CMakeLists.txt b/apps/ggml/halide/CMakeLists.txt new file mode 100644 index 000000000000..8adce0e4b9ea --- /dev/null +++ b/apps/ggml/halide/CMakeLists.txt @@ -0,0 +1,773 @@ +add_halide_generator( + quants.generator + SOURCES + f16_generators.cpp + bf16_generators.cpp + repack_quantize_mat_generators.cpp + repack_matmul_generator.cpp + symmetric_quant_generators.cpp + symmetric_vec_dot_generator.cpp + lookup_table_quant_generators.cpp + lookup_table_vec_dot_generator.cpp + k_quant_vec_dot_generator.cpp + k_quant_generators.cpp +) + +# Q4_0's and Q8_0's quantize/dequantize/vec_dot kernels are GENERATOR_ARGS +# instantiations of the generic, reusable Approximation-based +# symmetric_quantize/symmetric_dequantize/symmetric_vec_dot generators (see +# symmetric_quant_generators.cpp/symmetric_vec_dot_generator.cpp and +# quant_components.h) -- not their own per-format C++ Generator classes. See +# quant_components.h's RoundingMode/ScaleAnchor for what these params mean. +add_halide_library( + q4_0_quantize + FROM quants.generator + GENERATOR symmetric_quantize + PARAMS + block_size=32 qmax=8 code_bits=4 rounding=truncate_half_up_with_offset anchor=extreme_signed +) +add_halide_library( + q4_0_dequantize + FROM quants.generator + GENERATOR symmetric_dequantize + PARAMS + block_size=32 qmax=8 code_bits=4 rounding=truncate_half_up_with_offset anchor=extreme_signed +) +add_halide_library( + q4_0_vec_dot + FROM quants.generator + GENERATOR symmetric_vec_dot + FEATURES no_asserts no_bounds_query + PARAMS + w_kind=symmetric block_size=32 w_qmax=8 w_code_bits=4 w_rounding=truncate_half_up_with_offset + w_anchor=extreme_signed a_kind=q8_0 a_qmax=127 +) +# Q4_1 is affine (min+scale, not symmetric): quant_components.h's +# AffineQuantize + NibblePack, matching block_q4_1's {fp16 d; fp16 m; qs[16];}. +add_halide_library( + q4_1_quantize + FROM quants.generator + GENERATOR symmetric_quantize + PARAMS kind=affine block_size=32 levels=15 code_bits=4 affine_rounding=clamped_int8 +) +add_halide_library( + q4_1_dequantize + FROM quants.generator + GENERATOR symmetric_dequantize + PARAMS kind=affine block_size=32 levels=15 code_bits=4 affine_rounding=clamped_int8 +) +add_halide_library( + q4_1_vec_dot + FROM quants.generator + GENERATOR symmetric_vec_dot + FEATURES no_asserts no_bounds_query + PARAMS + w_kind=affine block_size=32 w_levels=15 w_code_bits=4 w_affine_rounding=clamped_int8 a_kind=q8_1 + a_qmax=127 +) +# Q5_0 is symmetric like Q4_0, but 5-bit: quant_components.h's +# SymmetricAffineQuantize + FiveBitPack, matching block_q5_0's +# {fp16 d; qh[4]; qs[16];}. +add_halide_library( + q5_0_quantize + FROM quants.generator + GENERATOR symmetric_quantize + PARAMS kind=symmetric_5bit block_size=32 qmax=16 +) +add_halide_library( + q5_0_dequantize + FROM quants.generator + GENERATOR symmetric_dequantize + PARAMS kind=symmetric_5bit block_size=32 qmax=16 +) +add_halide_library( + q5_0_vec_dot + FROM quants.generator + GENERATOR symmetric_vec_dot + PARAMS w_kind=symmetric_5bit block_size=32 w_qmax=16 a_kind=q8_0 a_qmax=127 + FEATURES no_asserts no_bounds_query +) +# Q5_1 is affine like Q4_1, but 5-bit: quant_components.h's AffineQuantize + +# FiveBitPack, matching block_q5_1's {fp16 d; fp16 m; qh[4]; qs[16];}. +add_halide_library( + q5_1_quantize + FROM quants.generator + GENERATOR symmetric_quantize + PARAMS kind=affine_5bit block_size=32 levels=31 affine_rounding=unclamped_uint8 +) +add_halide_library( + q5_1_dequantize + FROM quants.generator + GENERATOR symmetric_dequantize + PARAMS kind=affine_5bit block_size=32 levels=31 affine_rounding=unclamped_uint8 +) +add_halide_library( + q5_1_vec_dot + FROM quants.generator + GENERATOR symmetric_vec_dot + PARAMS + w_kind=affine_5bit block_size=32 w_levels=31 w_affine_rounding=unclamped_uint8 a_kind=q8_1 + a_qmax=127 + FEATURES no_asserts no_bounds_query +) +add_halide_library( + q8_0_quantize + FROM quants.generator + GENERATOR symmetric_quantize + PARAMS block_size=32 qmax=127 code_bits=8 rounding=nearest anchor=abs_max +) +add_halide_library( + q8_0_dequantize + FROM quants.generator + GENERATOR symmetric_dequantize + PARAMS block_size=32 qmax=127 code_bits=8 rounding=nearest anchor=abs_max +) +add_halide_library( + q8_0_vec_dot + FROM quants.generator + GENERATOR symmetric_vec_dot + FEATURES no_asserts no_bounds_query + PARAMS + w_kind=symmetric block_size=32 w_qmax=127 w_code_bits=8 w_rounding=nearest w_anchor=abs_max + a_kind=q8_0 a_qmax=127 +) +# Q8_1 is symmetric byte-packed like Q8_0, plus AppendCodeSum's derived 's' +# field, matching block_q8_1's {fp16 d; fp16 s; qs[32];}. Activation-only +# (no public to_float, see q8_1_generators.cpp's header comment -- it's +# gone now, but this scheme's decode() is still used by any vec_dot pairing +# against Q8_1), so there's no q8_1_dequantize library. +add_halide_library( + q8_1_quantize + FROM quants.generator + GENERATOR symmetric_quantize + PARAMS kind=symmetric_byte_sum block_size=32 qmax=127 +) +# Q8_K's quantize kernel (activation-only, no dequantize -- see +# q8_k_generators.cpp) is a GENERATOR_ARGS instantiation of the generic, +# reusable Approximation-based symmetric_quantize generator (see +# symmetric_quant_generators.cpp and quant_components.h's +# AppendGroupSumsInt16/F32Pack/RoundingMode::NearestEvenClampedHigh/ +# ScaleAnchor::ExtremeSignedValueTwoStep). +add_halide_library( + q8_k_quantize + FROM quants.generator + GENERATOR symmetric_quantize + PARAMS kind=q8k block_size=256 qmax=127 +) +# Q2_K's quantize/dequantize kernels are GENERATOR_ARGS instantiations of +# the generic, reusable Approximation-based k_quant_quantize/ +# k_quant_dequantize generators (see k_quant_generators.cpp and +# quant_components.h's KQuantDequantize/NibblePairPack/PlanarBitPack). +add_halide_library( + q2_k_quantize + FROM quants.generator + GENERATOR k_quant_quantize + PARAMS family=q2_k +) +add_halide_library( + q2_k_dequantize + FROM quants.generator + GENERATOR k_quant_dequantize + PARAMS family=q2_k +) +add_halide_library(q2_k_vec_dot FROM quants.generator GENERATOR k_quant_vec_dot PARAMS family=q2_k) +# Q6_K's quantize/dequantize kernels are GENERATOR_ARGS instantiations of +# the generic, reusable Approximation-based k_quant_quantize/ +# k_quant_dequantize generators (see k_quant_generators.cpp and +# quant_components.h's KQuantDequantize/CombinedBitsCode/BytePack). +add_halide_library( + q6_k_quantize + FROM quants.generator + GENERATOR k_quant_quantize + PARAMS family=q6_k +) +add_halide_library( + q6_k_dequantize + FROM quants.generator + GENERATOR k_quant_dequantize + PARAMS family=q6_k +) +add_halide_library(q6_k_vec_dot FROM quants.generator GENERATOR k_quant_vec_dot PARAMS family=q6_k) +# Q4_K's quantize/dequantize kernels are GENERATOR_ARGS instantiations of +# the generic, reusable Approximation-based k_quant_quantize/ +# k_quant_dequantize generators (see k_quant_generators.cpp and +# quant_components.h's KQuantDequantize/K4ScaleMinPack/PlanarBitPack). +add_halide_library( + q4_k_quantize + FROM quants.generator + GENERATOR k_quant_quantize + PARAMS family=q4_k +) +add_halide_library( + q4_k_dequantize + FROM quants.generator + GENERATOR k_quant_dequantize + PARAMS family=q4_k +) +add_halide_library(q4_k_vec_dot FROM quants.generator GENERATOR k_quant_vec_dot PARAMS family=q4_k) +# Q5_K's quantize/dequantize kernels are GENERATOR_ARGS instantiations of +# the generic, reusable Approximation-based k_quant_quantize/ +# k_quant_dequantize generators (see k_quant_generators.cpp and +# quant_components.h's KQuantDequantize/K4ScaleMinPack/CombinedBitsCode). +add_halide_library( + q5_k_quantize + FROM quants.generator + GENERATOR k_quant_quantize + PARAMS family=q5_k +) +add_halide_library( + q5_k_dequantize + FROM quants.generator + GENERATOR k_quant_dequantize + PARAMS family=q5_k +) +add_halide_library(q5_k_vec_dot FROM quants.generator GENERATOR k_quant_vec_dot PARAMS family=q5_k) +# Q3_K's quantize/dequantize kernels are GENERATOR_ARGS instantiations of +# the generic, reusable Approximation-based k_quant_quantize/ +# k_quant_dequantize generators (see k_quant_generators.cpp and +# quant_components.h's KQuantDequantize/Q3KScalePack/CombinedBitsCode). +add_halide_library( + q3_k_quantize + FROM quants.generator + GENERATOR k_quant_quantize + PARAMS family=q3_k +) +add_halide_library( + q3_k_dequantize + FROM quants.generator + GENERATOR k_quant_dequantize + PARAMS family=q3_k +) +add_halide_library(q3_k_vec_dot FROM quants.generator GENERATOR k_quant_vec_dot PARAMS family=q3_k) +# Q1_0 is symmetric with a mean-abs scale and sign-only (1-bit) codes: +# quant_components.h's SymmetricAffineQuantize (ScaleAnchor::MeanAbs, +# RoundingMode::SignOnly) + BitPack, matching block_q1_0's {fp16 d; qs[16];}. +add_halide_library( + q1_0_quantize + FROM quants.generator + GENERATOR symmetric_quantize + PARAMS block_size=128 qmax=1 code_bits=1 rounding=sign_only anchor=mean_abs +) +add_halide_library( + q1_0_dequantize + FROM quants.generator + GENERATOR symmetric_dequantize + PARAMS block_size=128 qmax=1 code_bits=1 rounding=sign_only anchor=mean_abs +) +add_halide_library( + q1_0_vec_dot + FROM quants.generator + GENERATOR symmetric_vec_dot + PARAMS + w_kind=symmetric block_size=128 w_qmax=1 w_code_bits=1 w_rounding=sign_only w_anchor=mean_abs + a_kind=q8_0 a_qmax=127 +) +# MXFP4/IQ4_NL's quantize/dequantize kernels are GENERATOR_ARGS +# instantiations of the generic, reusable Approximation-based +# lookup_table_quantize/lookup_table_dequantize generators (see +# lookup_table_quant_generators.cpp and quant_components.h's +# LookupTableQuantize/E8M0Pack) -- not their own per-format C++ Generator +# classes. vec_dot is still its own hand-rolled, unscheduled Generator (see +# mxfp4_generators.cpp/iq4_nl_generators.cpp), matching Q4_0/Q8_0's own +# treatment before they got the symmetric_vec_dot generic generator. +add_halide_library( + mxfp4_quantize + FROM quants.generator + GENERATOR lookup_table_quantize + PARAMS family=mxfp4 +) +add_halide_library( + mxfp4_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=mxfp4 +) +add_halide_library( + mxfp4_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=mxfp4 +) +# NVFP4's quantize/dequantize kernels are GENERATOR_ARGS instantiations of +# the generic, reusable Approximation-based lookup_table_quantize/ +# lookup_table_dequantize generators (see lookup_table_quant_generators.cpp +# and quant_components.h's LookupTableQuantize's num_scales/UE4M3Pack). +add_halide_library( + nvfp4_quantize + FROM quants.generator + GENERATOR lookup_table_quantize + PARAMS family=nvfp4 +) +add_halide_library( + nvfp4_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=nvfp4 +) +add_halide_library( + nvfp4_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=nvfp4 +) +add_halide_library( + iq4_nl_quantize + FROM quants.generator + GENERATOR lookup_table_quantize + PARAMS family=iq4_nl +) +add_halide_library( + iq4_nl_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=iq4_nl +) +add_halide_library( + iq4_nl_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=iq4_nl +) +# IQ4_XS's quantize/dequantize kernels are GENERATOR_ARGS instantiations of +# the generic, reusable Approximation-based lookup_table_quantize/ +# lookup_table_dequantize generators (see lookup_table_quant_generators.cpp +# and quant_components.h's IQ4XSDequantize). +add_halide_library( + iq4_xs_quantize + FROM quants.generator + GENERATOR lookup_table_quantize + PARAMS family=iq4_xs +) +add_halide_library( + iq4_xs_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=iq4_xs +) +add_halide_library( + iq4_xs_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=iq4_xs +) +# TQ1_0's quantize/dequantize kernels are GENERATOR_ARGS instantiations of +# the generic, reusable Approximation-based lookup_table_quantize/ +# lookup_table_dequantize generators (see lookup_table_quant_generators.cpp +# and quant_components.h's LookupTableQuantize/TritPack). +add_halide_library( + tq1_0_quantize + FROM quants.generator + GENERATOR lookup_table_quantize + PARAMS family=tq1_0 +) +add_halide_library( + tq1_0_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=tq1_0 +) +add_halide_library( + tq1_0_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=tq1_0 +) +# TQ2_0's quantize/dequantize kernels are GENERATOR_ARGS instantiations of +# the generic, reusable Approximation-based lookup_table_quantize/ +# lookup_table_dequantize generators (see lookup_table_quant_generators.cpp +# and quant_components.h's LookupTableQuantize/PlanarBitPack). +add_halide_library( + tq2_0_quantize + FROM quants.generator + GENERATOR lookup_table_quantize + PARAMS family=tq2_0 +) +add_halide_library( + tq2_0_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=tq2_0 +) +add_halide_library( + tq2_0_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=tq2_0 +) +add_halide_library( + iq2_xxs_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=iq2_xxs +) +add_halide_library( + iq2_xxs_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=iq2_xxs +) +add_halide_library( + iq2_xs_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=iq2_xs +) +add_halide_library( + iq2_xs_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=iq2_xs +) +# IQ2_S/IQ3_XXS/IQ3_S's quantize/dequantize kernels are GENERATOR_ARGS +# instantiations of the generic, reusable Approximation-based +# lookup_table_quantize/lookup_table_dequantize generators (see +# lookup_table_quant_generators.cpp and quant_components.h's +# IQ2SGridDequantize/IQ3XXSGridDequantize/IQ3SGridDequantize). +add_halide_library( + iq2_s_quantize + FROM quants.generator + GENERATOR lookup_table_quantize + PARAMS family=iq2_s +) +add_halide_library( + iq2_s_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=iq2_s +) +add_halide_library( + iq2_s_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=iq2_s +) +add_halide_library( + iq3_xxs_quantize + FROM quants.generator + GENERATOR lookup_table_quantize + PARAMS family=iq3_xxs +) +add_halide_library( + iq3_xxs_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=iq3_xxs +) +add_halide_library( + iq3_xxs_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=iq3_xxs +) +add_halide_library( + iq3_s_quantize + FROM quants.generator + GENERATOR lookup_table_quantize + PARAMS family=iq3_s +) +add_halide_library( + iq3_s_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=iq3_s +) +add_halide_library( + iq3_s_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=iq3_s +) +add_halide_library( + iq1_s_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=iq1_s +) +add_halide_library( + iq1_s_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=iq1_s +) +add_halide_library( + iq1_m_dequantize + FROM quants.generator + GENERATOR lookup_table_dequantize + PARAMS family=iq1_m +) +add_halide_library( + iq1_m_vec_dot + FROM quants.generator + GENERATOR lookup_table_vec_dot + PARAMS family=iq1_m +) +add_halide_library(f16_quantize FROM quants.generator GENERATOR f16_quantize) +add_halide_library(f16_dequantize FROM quants.generator GENERATOR f16_dequantize) +add_halide_library(bf16_quantize FROM quants.generator GENERATOR bf16_quantize) +add_halide_library(bf16_dequantize FROM quants.generator GENERATOR bf16_dequantize) +add_halide_library(q8_0_4x4_quantize_mat FROM quants.generator GENERATOR q8_0_4x4_quantize_mat) +add_halide_library(q8_0_4x8_quantize_mat FROM quants.generator GENERATOR q8_0_4x8_quantize_mat) +add_halide_library(q8_k_4x4_quantize_mat FROM quants.generator GENERATOR q8_k_4x4_quantize_mat) +add_halide_library(q8_k_4x8_quantize_mat FROM quants.generator GENERATOR q8_k_4x8_quantize_mat) + +add_halide_library( + q4_0_4x4_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q4_0 n_cols=4 blocklen=4 +) +add_halide_library( + q4_0_4x8_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q4_0 n_cols=4 blocklen=8 +) +add_halide_library( + q4_0_8x8_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q4_0 n_cols=8 blocklen=8 +) +add_halide_library( + q8_0_4x4_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q8_0 n_cols=4 blocklen=4 +) +add_halide_library( + q8_0_4x8_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q8_0 n_cols=4 blocklen=8 +) +add_halide_library( + q4_0_4x4_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q4_0 n_cols=4 blocklen=4 +) +add_halide_library( + q4_0_4x8_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q4_0 n_cols=4 blocklen=8 +) +add_halide_library( + q4_0_8x8_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q4_0 n_cols=8 blocklen=8 +) +add_halide_library( + q8_0_4x4_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q8_0 n_cols=4 blocklen=4 +) +add_halide_library( + q8_0_4x8_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q8_0 n_cols=4 blocklen=8 +) +add_halide_library( + iq4_nl_4x4_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=iq4_nl n_cols=4 blocklen=4 +) +add_halide_library( + iq4_nl_8x8_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=iq4_nl n_cols=8 blocklen=8 +) +add_halide_library( + mxfp4_4x4_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=mxfp4 n_cols=4 blocklen=4 +) +add_halide_library( + mxfp4_8x8_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=mxfp4 n_cols=8 blocklen=8 +) +add_halide_library( + iq4_nl_4x4_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=iq4_nl n_cols=4 blocklen=4 +) +add_halide_library( + iq4_nl_8x8_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=iq4_nl n_cols=8 blocklen=8 +) +add_halide_library( + mxfp4_4x4_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=mxfp4 n_cols=4 blocklen=4 +) +add_halide_library( + mxfp4_8x8_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=mxfp4 n_cols=8 blocklen=8 +) +add_halide_library( + q4_k_8x4_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q4_k n_cols=8 blocklen=4 +) +add_halide_library( + q4_k_8x8_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q4_k n_cols=8 blocklen=8 +) +add_halide_library( + q4_k_8x4_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q4_k n_cols=8 blocklen=4 +) +add_halide_library( + q4_k_8x8_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q4_k n_cols=8 blocklen=8 +) +add_halide_library( + q5_k_8x4_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q5_k n_cols=8 blocklen=4 +) +add_halide_library( + q5_k_8x8_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q5_k n_cols=8 blocklen=8 +) +add_halide_library( + q5_k_8x4_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q5_k n_cols=8 blocklen=4 +) +add_halide_library( + q5_k_8x8_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q5_k n_cols=8 blocklen=8 +) +add_halide_library( + q6_k_8x4_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q6_k n_cols=8 blocklen=4 +) +add_halide_library( + q6_k_8x8_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q6_k n_cols=8 blocklen=8 +) +add_halide_library( + q6_k_8x4_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q6_k n_cols=8 blocklen=4 +) +add_halide_library( + q6_k_8x8_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q6_k n_cols=8 blocklen=8 +) +add_halide_library( + q2_k_8x8_gemv + FROM quants.generator + GENERATOR repack_gemv + PARAMS family=q2_k n_cols=8 blocklen=8 +) +add_halide_library( + q2_k_8x8_gemm + FROM quants.generator + GENERATOR repack_gemm + PARAMS family=q2_k n_cols=8 blocklen=8 +) + +# ggml_extern_quantize.cpp implements the extern-stage C functions that +# several of the quantize pipelines above call out to -- see that file for +# why. It's the one source in this library that depends on GGML, hence the +# extra ggml::ggml link below (every other source here is GGML-independent). +add_library(ggml_quants_halide ggml_quants.cpp ggml_extern_quantize.cpp) +target_link_libraries( + ggml_quants_halide + PUBLIC + q4_0_quantize q4_0_dequantize q4_0_vec_dot q4_1_quantize q4_1_dequantize q4_1_vec_dot + q5_0_quantize q5_0_dequantize q5_0_vec_dot q5_1_quantize q5_1_dequantize q5_1_vec_dot + q8_0_quantize q8_0_dequantize q8_0_vec_dot q8_1_quantize q8_k_quantize q2_k_quantize + q2_k_dequantize q2_k_vec_dot q6_k_quantize q6_k_dequantize q6_k_vec_dot q4_k_quantize + q4_k_dequantize q4_k_vec_dot q5_k_quantize q5_k_dequantize q5_k_vec_dot q3_k_quantize + q3_k_dequantize q3_k_vec_dot q1_0_quantize q1_0_dequantize q1_0_vec_dot mxfp4_quantize + mxfp4_dequantize mxfp4_vec_dot nvfp4_quantize nvfp4_dequantize nvfp4_vec_dot iq4_nl_quantize + iq4_nl_dequantize iq4_nl_vec_dot iq4_xs_quantize iq4_xs_dequantize iq4_xs_vec_dot tq1_0_quantize + tq1_0_dequantize tq1_0_vec_dot tq2_0_quantize tq2_0_dequantize tq2_0_vec_dot iq2_xxs_dequantize + iq2_xxs_vec_dot iq2_xs_dequantize iq2_xs_vec_dot iq2_s_quantize iq2_s_dequantize iq2_s_vec_dot + iq3_xxs_quantize iq3_xxs_dequantize iq3_xxs_vec_dot iq3_s_quantize iq3_s_dequantize + iq3_s_vec_dot iq1_s_dequantize iq1_s_vec_dot iq1_m_dequantize iq1_m_vec_dot f16_quantize + f16_dequantize bf16_quantize bf16_dequantize q8_0_4x4_quantize_mat q8_0_4x8_quantize_mat + q8_k_4x4_quantize_mat q8_k_4x8_quantize_mat q4_0_4x4_gemv q4_0_4x8_gemv q4_0_8x8_gemv + q8_0_4x4_gemv q8_0_4x8_gemv q4_0_4x4_gemm q4_0_4x8_gemm q4_0_8x8_gemm q8_0_4x4_gemm + q8_0_4x8_gemm iq4_nl_4x4_gemv iq4_nl_8x8_gemv mxfp4_4x4_gemv mxfp4_8x8_gemv iq4_nl_4x4_gemm + iq4_nl_8x8_gemm mxfp4_4x4_gemm mxfp4_8x8_gemm q4_k_8x4_gemv q4_k_8x8_gemv q4_k_8x4_gemm + q4_k_8x8_gemm q5_k_8x4_gemv q5_k_8x8_gemv q5_k_8x4_gemm q5_k_8x8_gemm q6_k_8x4_gemv + q6_k_8x8_gemv q6_k_8x4_gemm q6_k_8x8_gemm q2_k_8x8_gemv q2_k_8x8_gemm ggml::ggml +) +target_include_directories(ggml_quants_halide PUBLIC ${CMAKE_CURRENT_SOURCE_DIR}) + +# Standalone correctness checks against GGML's own reference, run before (and +# independently of) wiring these into kernel-bench as providers. +foreach (t IN ITEMS + q4_0 + q4_1 + q5_0 + q5_1 + q8_0 + q8_1 + q8_k + q2_k + q6_k + q4_k + q5_k + q3_k + q1_0 + mxfp4 + nvfp4 + iq4_nl + iq4_xs + tq1_0 + tq2_0 + iq2_xxs + iq2_xs + iq2_s + iq3_xxs + iq3_s + iq1_s + iq1_m + f16 + bf16 +) + add_executable(test_${t} test_${t}.cpp) + target_link_libraries(test_${t} PRIVATE ggml_quants_halide ggml::ggml) + target_include_directories(test_${t} PRIVATE ${CMAKE_CURRENT_SOURCE_DIR}/../src) + add_test(NAME ${t}_roundtrip COMMAND test_${t}) + set_tests_properties(${t}_roundtrip PROPERTIES PASS_REGULAR_EXPRESSION "Success!") +endforeach () diff --git a/apps/ggml/halide/bf16_generators.cpp b/apps/ggml/halide/bf16_generators.cpp new file mode 100644 index 000000000000..bab89572cb5f --- /dev/null +++ b/apps/ggml/halide/bf16_generators.cpp @@ -0,0 +1,58 @@ +// From-scratch Halide reimplementation of GGML's BF16 quantize/dequantize +// "kernels" (see src/ggml-impl.h: ggml_compute_fp32_to_bf16 / +// ggml_compute_bf16_to_fp32 upstream, as of GGML v0.15.3). Like F16, BF16 +// isn't really a quantized format -- it's a 1-element/block cast, no +// header/payload split -- so this is just Halide's native bfloat16_t cast +// in both directions. +// +// GGML's bf16 encode is round-to-nearest-even truncation of the top 16 bits +// of the IEEE binary32 representation (with NaNs forced quiet); decode is +// a plain `bits << 16` reinterpretation. Halide's cast/cast +// compile to the same IEEE-mandated round-to-nearest-even truncation and +// zero-extension, so this matches bit-for-bit for all finite inputs (the +// only divergence possible is NaN payload/quieting, which the synthetic +// benchmark/test data never produces). +// +// This is intentionally unscheduled -- scheduling for performance is a +// later step. + +#include "Halide.h" + +using namespace Halide; + +namespace { + +class BF16DequantizeGenerator : public Generator { +public: + // Raw bf16 bit patterns, one uint16 per element (block size 1). + Input> x_{"x"}; + Output> y_{"y"}; + + void generate() { + Var i("i"); + y_(i) = cast(reinterpret(x_(i))); + + x_.dim(0).set_min(0); + y_.dim(0).set_min(0); + } +}; + +class BF16QuantizeGenerator : public Generator { +public: + Input> x_{"x"}; + // Raw bf16 bit patterns, one uint16 per element (block size 1). + Output> y_{"y"}; + + void generate() { + Var i("i"); + y_(i) = reinterpret(cast(x_(i))); + + x_.dim(0).set_min(0); + y_.dim(0).set_min(0); + } +}; + +} // namespace + +HALIDE_REGISTER_GENERATOR(BF16DequantizeGenerator, bf16_dequantize) +HALIDE_REGISTER_GENERATOR(BF16QuantizeGenerator, bf16_quantize) diff --git a/apps/ggml/halide/codec_generator_base.h b/apps/ggml/halide/codec_generator_base.h new file mode 100644 index 000000000000..54293d376ee7 --- /dev/null +++ b/apps/ggml/halide/codec_generator_base.h @@ -0,0 +1,126 @@ +#pragma once + +// Shared configure()/generate() scaffolding for every *_quant_generators.cpp +// file's Direction-templated Generator (SymmetricCodecGenerator, +// LookupTableCodecGenerator, KQuantCodecGenerator): all three build the +// exact same "real ImageParam -> approximate_by -> sever -> adopt +// one half as a port" pipeline in configure(), differing only in how their +// SchemeAndBytes gets built. This factors that shared body out via CRTP +// (Derived::build_scheme()) -- the same static-polymorphism idiom +// Halide::Generator itself already uses (see its own `T` template +// parameter), not a virtual method: the concrete type is always known at +// compile time, so there's no reason to pay for a vtable. Confirmed safe to +// insert as a base class between a leaf Generator and Halide::Generator: +// GeneratorParam/Input/Output discovery is address-range-based (see +// Generator.cpp's ObjectInstanceRegistry::register_instance/ +// instances_in_range), not declaration-order or hierarchy-position based, +// so it doesn't matter which class in the chain declares them. +// +// Usage: +// class FooCodecGenerator : public CodecGeneratorBase, dir> { +// public: +// GeneratorParam<...> whatever{...}; +// SchemeAndBytes build_scheme() const { return ::build_scheme(whatever); } +// }; + +#include "Halide.h" + +#include "quant_components.h" + +namespace ggml_halide { + +enum class Direction { Quantize, + Dequantize }; + +// SchemeAndBytes itself now lives in quant_components.h (its `scheme` is +// held as a polymorphic owning handle -- a single leaf, a Compose, or a +// TrustedInverse, whichever the format is; see the make_*() factories +// there) -- moved there so those factories can return it directly instead +// of every Generator switch hand-summing a byte count alongside a bare +// scheme. + +template +class CodecGeneratorBase : public Halide::Generator { +public: + void configure() { + using namespace Halide; + SchemeAndBytes sb = static_cast(this)->build_scheme(); + + // A structured scheme's encoded form is a first-class 1-D Type::Struct + // block (one struct per block index); an unported one is the flat 2-D + // (byte, blk) UInt(8) buffer. block_type.bytes() is the on-disk width in + // the struct case -- no separately-threaded block_bytes needed. + const bool structured = sb.block_type.is_struct(); + + // The "obvious" identity: a real ImageParam (never a placeholder -- + // that's what lets *both* directions share this one call below) + // flowing through unchanged. + Var x("x"); + ImageParam input(Float(32), 1, "x"); + Func identity("y"); + identity(x) = input(x); + + ApproximationResult r = Func(input).approximate_by(sb.scheme, {identity}); + // Materialize the reductions and every stage-boundary Func; other pure + // Funcs inside a stage stay inline. The Q5 struct layout's packed + // high-bit word is the one non-port pure Func that is also + // materialized (found by name, like vec_dot_generator_base.h). + for (Func h : r.intermediates) { + if (h.has_update_definition() || r.is_stage_port(h) || + h.name() == "q5_struct_block_qh") { + h.compute_root(); + } + } + + // Bind sever() to a properly-named ImageParam of the packed + // block's shape up front, instead of letting it mint one named after + // whatever internal Func produced r.encoded[0]. Only Dequantize below + // adopts it as a port (named "blocks_in" rather than reusing Quantize's + // output name "blocks_out"); the two are never both real ports at once, + // but both objects always exist. + ImageParam blocks_in = structured ? ImageParam(sb.block_type, 1, "blocks_in") : ImageParam(UInt(8), 2, "blocks_in"); + + // Severs `identity` from `input`/encode() entirely: `q.offline` + // recomputes r.encoded (quantize) from `input`, while `identity` + // (post-severance) instead reads from `blocks_in` (dequantize). + SeverResult q = Pipeline({identity}).sever(r.encoded, {blocks_in}); + + if constexpr (dir == Direction::Quantize) { + input.dim(0).set_min(0); + + // A thin renamed passthrough gives the compiled Output a clean name + // (the way `blocks_in` did for the Input side); Halide inlines it. + Func blocks_out("blocks_out"); + Var byte("byte"), blk("blk"); + if (structured) { + blocks_out(blk) = q.offline.outputs()[0](blk); + blocks_out.output_buffer().dim(0).set_min(0); + } else { + blocks_out(byte, blk) = q.offline.outputs()[0](byte, blk); + blocks_out.output_buffer().dim(0).set_bounds(0, sb.block_bytes); + blocks_out.output_buffer().dim(1).set_min(0); + } + + this->add_input(input); + this->add_output(blocks_out); + } else { + if (structured) { + blocks_in.dim(0).set_min(0); + } else { + blocks_in.dim(0).set_bounds(0, sb.block_bytes); + blocks_in.dim(1).set_min(0); + } + identity.output_buffer().dim(0).set_min(0); + + this->add_input(blocks_in); + this->add_output(identity); + } + } + + void generate() { + // Nothing left to do: configure() already built (and, via + // add_input/add_output, wired up) the whole pipeline. + } +}; + +} // namespace ggml_halide diff --git a/apps/ggml/halide/f16_generators.cpp b/apps/ggml/halide/f16_generators.cpp new file mode 100644 index 000000000000..9e5431df20f6 --- /dev/null +++ b/apps/ggml/halide/f16_generators.cpp @@ -0,0 +1,55 @@ +// From-scratch Halide reimplementation of GGML's F16 quantize/dequantize +// "kernels" (see src/ggml.c: ggml_fp32_to_fp16_row / ggml_fp16_to_fp32_row +// upstream, as of GGML v0.15.3). F16 isn't really a quantized format -- it's +// a 1-element/block plain IEEE-754 binary16 cast, no header/payload split at +// all -- so unlike every other type here, this is just Halide's native +// float16_t cast in both directions. +// +// GGML's own conversion (GGML_COMPUTE_FP32_TO_FP16 / _FP16_TO_FP32 in +// src/ggml-impl.h) is a correctly-rounded (round-to-nearest-even) software +// IEEE binary16 <-> binary32 conversion; Halide's cast/cast +// compiles to the same IEEE-mandated conversion, so this matches bit-for-bit. +// +// This is intentionally unscheduled -- scheduling for performance is a +// later step. + +#include "Halide.h" + +using namespace Halide; + +namespace { + +class F16DequantizeGenerator : public Generator { +public: + // Raw fp16 bit patterns, one uint16 per element (block size 1). + Input> x_{"x"}; + Output> y_{"y"}; + + void generate() { + Var i("i"); + y_(i) = cast(reinterpret(x_(i))); + + x_.dim(0).set_min(0); + y_.dim(0).set_min(0); + } +}; + +class F16QuantizeGenerator : public Generator { +public: + Input> x_{"x"}; + // Raw fp16 bit patterns, one uint16 per element (block size 1). + Output> y_{"y"}; + + void generate() { + Var i("i"); + y_(i) = reinterpret(cast(x_(i))); + + x_.dim(0).set_min(0); + y_.dim(0).set_min(0); + } +}; + +} // namespace + +HALIDE_REGISTER_GENERATOR(F16DequantizeGenerator, f16_dequantize) +HALIDE_REGISTER_GENERATOR(F16QuantizeGenerator, f16_quantize) diff --git a/apps/ggml/halide/ggml_extern_quantize.cpp b/apps/ggml/halide/ggml_extern_quantize.cpp new file mode 100644 index 000000000000..519b8fd29878 --- /dev/null +++ b/apps/ggml/halide/ggml_extern_quantize.cpp @@ -0,0 +1,116 @@ +// Extern-stage scaffolding for K-quant quantize kernels. +// +// GGML's reference quantizer for the K-quant super-block formats (Q2_K, +// Q3_K, Q4_K, Q5_K, Q6_K) isn't a closed-form scale computation like every +// other type in this directory -- it runs an iterative, per-sub-block +// error-minimizing search over ~19 candidate scale factors (see +// src/ggml-quants.c: make_qx_quants / make_qkx1_quants / make_qkx2_quants / +// make_q3_quants). Porting that search to Halide is deferred; per the +// project's current phase, this sets up the Halide extern-stage plumbing +// now (a Func whose realization is computed by an external C function) and +// simply calls out to GGML's own public from_float_ref for the actual +// computation. Dequantize (a pure unpacking operation, no search) is +// implemented natively in Halide for these types -- see qX_k_generators.cpp. +// +// This is the one file in halide/ that depends on GGML's public API -- +// every generator's *body* stays GGML-independent, but this scaffold +// deliberately borrows GGML's own reference computation for now, to be +// replaced with a from-scratch Halide search later. +// +// Extern-stage ABI: a plain C function taking one halide_buffer_t* per +// Func argument/output, returning 0 on success. Halide calls it twice per +// realization: once in "bounds query" mode (host pointers null, dimensions +// need to be filled in based on the output's already-concrete request) and +// once for real (host pointers valid, actually compute the data). See +// test/correctness/extern_bounds_inference.cpp for the reference pattern. + +#include + +#include + +namespace { + +int quantize_via_ggml_reference(ggml_type type, halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + if (x_buf->is_bounds_query()) { + // out_buf already carries the concrete requested region (its dim[1] + // extent is the number of blocks); the input row needed is exactly + // that many blocks' worth of elements. + const int64_t nb = out_buf->dim[1].extent; + x_buf->dim[0].min = 0; + x_buf->dim[0].extent = static_cast(nb * ggml_blck_size(type)); + return 0; + } + + const float *x = reinterpret_cast(x_buf->host); + void *y = reinterpret_cast(out_buf->host); + const int64_t k = x_buf->dim[0].extent; + + ggml_get_type_traits(type)->from_float_ref(x, y, k); + return 0; +} + +} // namespace + +extern "C" int q2_k_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_Q2_K, x_buf, out_buf); +} + +extern "C" int q3_k_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_Q3_K, x_buf, out_buf); +} + +extern "C" int q4_k_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_Q4_K, x_buf, out_buf); +} + +extern "C" int q5_k_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_Q5_K, x_buf, out_buf); +} + +extern "C" int q6_k_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_Q6_K, x_buf, out_buf); +} + +// The remaining types below use this same scaffolding for different +// reasons than the K-quants: MXFP4/NVFP4 derive their scale via a +// transcendental (log2) or rounding-sensitive fixed-point float format not +// guaranteed to be bit-reproducible from scratch; IQ4_NL/IQ4_XS run a +// per-block/sub-block nearest-codeword search with scale refinement; TQ1_0/ +// TQ2_0's byte-packing, while closed-form, is fiddly to unroll in Halide's +// functional style. All are deferred the same way, for now. + +extern "C" int mxfp4_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_MXFP4, x_buf, out_buf); +} + +extern "C" int nvfp4_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_NVFP4, x_buf, out_buf); +} + +extern "C" int iq4_nl_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_IQ4_NL, x_buf, out_buf); +} + +extern "C" int iq4_xs_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_IQ4_XS, x_buf, out_buf); +} + +extern "C" int tq1_0_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_TQ1_0, x_buf, out_buf); +} + +extern "C" int tq2_0_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_TQ2_0, x_buf, out_buf); +} + +extern "C" int iq3_xxs_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_IQ3_XXS, x_buf, out_buf); +} + +extern "C" int iq3_s_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_IQ3_S, x_buf, out_buf); +} + +extern "C" int iq2_s_quantize_via_ggml(halide_buffer_t *x_buf, halide_buffer_t *out_buf) { + return quantize_via_ggml_reference(GGML_TYPE_IQ2_S, x_buf, out_buf); +} diff --git a/apps/ggml/halide/ggml_quants.cpp b/apps/ggml/halide/ggml_quants.cpp new file mode 100644 index 000000000000..18acc52c35aa --- /dev/null +++ b/apps/ggml/halide/ggml_quants.cpp @@ -0,0 +1,1783 @@ +#include "ggml_quants.h" + +#include + +#include "HalideBuffer.h" +#include "bf16_dequantize.h" +#include "bf16_quantize.h" +#include "f16_dequantize.h" +#include "f16_quantize.h" +#include "iq1_m_dequantize.h" +#include "iq1_m_vec_dot.h" +#include "iq1_s_dequantize.h" +#include "iq1_s_vec_dot.h" +#include "iq2_s_dequantize.h" +#include "iq2_s_quantize.h" +#include "iq2_s_vec_dot.h" +#include "iq2_xs_dequantize.h" +#include "iq2_xs_vec_dot.h" +#include "iq2_xxs_dequantize.h" +#include "iq2_xxs_vec_dot.h" +#include "iq3_s_dequantize.h" +#include "iq3_s_quantize.h" +#include "iq3_s_vec_dot.h" +#include "iq3_xxs_dequantize.h" +#include "iq3_xxs_quantize.h" +#include "iq3_xxs_vec_dot.h" +#include "iq4_nl_4x4_gemm.h" +#include "iq4_nl_4x4_gemv.h" +#include "iq4_nl_8x8_gemm.h" +#include "iq4_nl_8x8_gemv.h" +#include "iq4_nl_dequantize.h" +#include "iq4_nl_quantize.h" +#include "iq4_nl_vec_dot.h" +#include "iq4_xs_dequantize.h" +#include "iq4_xs_quantize.h" +#include "iq4_xs_vec_dot.h" +#include "mxfp4_4x4_gemm.h" +#include "mxfp4_4x4_gemv.h" +#include "mxfp4_8x8_gemm.h" +#include "mxfp4_8x8_gemv.h" +#include "mxfp4_dequantize.h" +#include "mxfp4_quantize.h" +#include "mxfp4_vec_dot.h" +#include "nvfp4_dequantize.h" +#include "nvfp4_quantize.h" +#include "nvfp4_vec_dot.h" +#include "q1_0_dequantize.h" +#include "q1_0_quantize.h" +#include "q1_0_vec_dot.h" +#include "q2_k_8x8_gemm.h" +#include "q2_k_8x8_gemv.h" +#include "q2_k_dequantize.h" +#include "q2_k_quantize.h" +#include "q2_k_vec_dot.h" +#include "q3_k_dequantize.h" +#include "q3_k_quantize.h" +#include "q3_k_vec_dot.h" +#include "q4_0_4x4_gemm.h" +#include "q4_0_4x4_gemv.h" +#include "q4_0_4x8_gemm.h" +#include "q4_0_4x8_gemv.h" +#include "q4_0_8x8_gemm.h" +#include "q4_0_8x8_gemv.h" +#include "q4_0_dequantize.h" +#include "q4_0_quantize.h" +#include "q4_0_vec_dot.h" +#include "q4_1_dequantize.h" +#include "q4_1_quantize.h" +#include "q4_1_vec_dot.h" +#include "q4_k_8x4_gemm.h" +#include "q4_k_8x4_gemv.h" +#include "q4_k_8x8_gemm.h" +#include "q4_k_8x8_gemv.h" +#include "q4_k_dequantize.h" +#include "q4_k_quantize.h" +#include "q4_k_vec_dot.h" +#include "q5_0_dequantize.h" +#include "q5_0_quantize.h" +#include "q5_0_vec_dot.h" +#include "q5_1_dequantize.h" +#include "q5_1_quantize.h" +#include "q5_1_vec_dot.h" +#include "q5_k_8x4_gemm.h" +#include "q5_k_8x4_gemv.h" +#include "q5_k_8x8_gemm.h" +#include "q5_k_8x8_gemv.h" +#include "q5_k_dequantize.h" +#include "q5_k_quantize.h" +#include "q5_k_vec_dot.h" +#include "q6_k_8x4_gemm.h" +#include "q6_k_8x4_gemv.h" +#include "q6_k_8x8_gemm.h" +#include "q6_k_8x8_gemv.h" +#include "q6_k_dequantize.h" +#include "q6_k_quantize.h" +#include "q6_k_vec_dot.h" +#include "q8_0_4x4_gemm.h" +#include "q8_0_4x4_gemv.h" +#include "q8_0_4x4_quantize_mat.h" +#include "q8_0_4x8_gemm.h" +#include "q8_0_4x8_gemv.h" +#include "q8_0_4x8_quantize_mat.h" +#include "q8_0_dequantize.h" +#include "q8_0_quantize.h" +#include "q8_0_vec_dot.h" +#include "q8_1_quantize.h" +#include "q8_k_4x4_quantize_mat.h" +#include "q8_k_4x8_quantize_mat.h" +#include "q8_k_quantize.h" +#include "tq1_0_dequantize.h" +#include "tq1_0_quantize.h" +#include "tq1_0_vec_dot.h" +#include "tq2_0_dequantize.h" +#include "tq2_0_quantize.h" +#include "tq2_0_vec_dot.h" + +using Halide::Runtime::Buffer; + +namespace { + +void check(int result, const char *what) { + if (result != 0) { + std::fprintf(stderr, "ggml_quants_halide: %s failed (%d)\n", what, result); + } +} + +// A 1-D Type::Struct block buffer over `nb` blocks of `block_bytes` each, +// wrapping GGML's raw packed bytes at `data`. The struct's ABI tag is +// {halide_type_struct, bits=8, reserved=block_bytes}; each element is one whole +// block, dim-0 stride 1 in struct units (see Halide::Type::to_abi()). This is +// how a Phase-3 struct-typed codec kernel receives GGML's byte layout unchanged. +Halide::Runtime::Buffer struct_block_buffer(const void *data, int nb, int block_bytes) { + halide_type_t ty(halide_type_struct, 8); + ty.reserved = static_cast(block_bytes); + halide_dimension_t shape[1] = {{0, nb, 1}}; + return Halide::Runtime::Buffer(ty, const_cast(data), 1, shape); +} + +// vec_dot is called once per output element of a matvec, so the argument +// marshalling is not amortized over anything: at the row lengths GGML actually +// uses, building three Halide::Runtime::Buffers costs about as much as the dot +// product itself. These kernels know their shapes exactly, so the vec_dot +// wrappers below fill a halide_buffer_t in place instead. (The quantize / +// dequantize / gemv wrappers stay on Buffer -- they run over whole rows or +// tiles, where the difference is noise.) +struct StackBuffer { + halide_buffer_t buf{}; + halide_dimension_t dims[2]{}; + + // Packed blocks as a 1-D Type::Struct buffer: one struct per block, with the + // block width in the type's `reserved` field (see Halide::Type::to_abi()). + halide_buffer_t *blocks_struct(const void *data, int nb, int block_bytes) { + buf.type = halide_type_t(halide_type_struct, 8); + buf.type.reserved = static_cast(block_bytes); + dims[0] = {0, nb, 1, 0}; + return init(data, 1); + } + + // Packed blocks as a 2-D (byte, block) UInt(8) buffer. + halide_buffer_t *blocks_bytes(const void *data, int nb, int block_bytes) { + buf.type = halide_type_t(halide_type_uint, 8); + dims[0] = {0, block_bytes, 1, 0}; + dims[1] = {0, nb, block_bytes, 0}; + return init(data, 2); + } + + // A gathered 1-D fp16 view of one field within each packed block -- e.g. + // Q8_1's stored `s` (scaled code sum) at byte_offset 2 of its 36-byte block. + // Zero-copy: the field repeats every block_bytes, so the fp16 stride is + // block_bytes/2. Used to sever a stored per-block quantity into a vec_dot. + halide_buffer_t *blocks_field_f16(const void *base, int nb, int byte_offset, int block_bytes) { + buf.type = halide_type_t(halide_type_float, 16); + dims[0] = {0, nb, block_bytes / 2, 0}; // stride in fp16 units + return init(static_cast(base) + byte_offset, 1); + } + + // The 0-D float32 result. + halide_buffer_t *scalar_f32(float *data) { + buf.type = halide_type_t(halide_type_float, 32); + return init(data, 0); + } + +private: + halide_buffer_t *init(const void *data, int dimensions) { + buf.host = const_cast(static_cast(data)); + buf.dimensions = dimensions; + buf.dim = dimensions ? dims : nullptr; + return &buf; + } +}; + +} // namespace + +extern "C" { + +// +// Q4_0 -- block size 32, 18 bytes/block (2 delta + 16 packed nibbles). +// + +void ggml_quants_halide_quantize_q4_0(const float *x, void *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 2 + kQK / 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + auto blocks = struct_block_buffer(y, nb, kBlockBytes); + check(q4_0_quantize(xb, blocks), "q4_0_quantize"); +} + +void ggml_quants_halide_dequantize_q4_0(const void *x, float *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 2 + kQK / 2; + const int32_t nb = static_cast(k / kQK); + auto blocks = struct_block_buffer(x, nb, kBlockBytes); + Buffer yb(y, static_cast(k)); + check(q4_0_dequantize(blocks, yb), "q4_0_dequantize"); +} + +void ggml_quants_halide_vec_dot_q4_0_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 32, kBlockBytesX = 2 + kQK / 2, kBlockBytesY = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + StackBuffer xb, yb, result; + check(q4_0_vec_dot(xb.blocks_struct(vx, nb, kBlockBytesX), // weight: struct-typed + yb.blocks_bytes(vy, nb, kBlockBytesY), + result.scalar_f32(s)), + "q4_0_vec_dot"); +} + +// +// Q4_1 -- block size 32, 20 bytes/block (2 delta + 2 min + 16 packed nibbles). +// + +void ggml_quants_halide_quantize_q4_1(const float *x, void *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 4 + kQK / 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(q4_1_quantize(xb, blocks), "q4_1_quantize"); +} + +void ggml_quants_halide_dequantize_q4_1(const void *x, float *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 4 + kQK / 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(q4_1_dequantize(blocks, yb), "q4_1_dequantize"); +} + +void ggml_quants_halide_vec_dot_q4_1_q8_1(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 32, kBlockBytesX = 4 + kQK / 2, kBlockBytesY = 4 + kQK; + const int32_t nb = static_cast(n / kQK); + StackBuffer xb, yb, sb, result; + check(q4_1_vec_dot(xb.blocks_bytes(vx, nb, kBlockBytesX), + yb.blocks_bytes(vy, nb, kBlockBytesY), + sb.blocks_field_f16(vy, nb, 2, kBlockBytesY), // Q8_1 stored `s` + result.scalar_f32(s)), + "q4_1_vec_dot"); +} + +// +// Q5_0 -- block size 32, 22 bytes/block (2 delta + 4 qh + 16 packed nibbles). +// + +void ggml_quants_halide_quantize_q5_0(const float *x, void *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 2 + 4 + kQK / 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + auto blocks = struct_block_buffer(y, nb, kBlockBytes); + check(q5_0_quantize(xb, blocks), "q5_0_quantize"); +} + +void ggml_quants_halide_dequantize_q5_0(const void *x, float *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 2 + 4 + kQK / 2; + const int32_t nb = static_cast(k / kQK); + auto blocks = struct_block_buffer(x, nb, kBlockBytes); + Buffer yb(y, static_cast(k)); + check(q5_0_dequantize(blocks, yb), "q5_0_dequantize"); +} + +void ggml_quants_halide_vec_dot_q5_0_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 32, kBlockBytesX = 2 + 4 + kQK / 2, kBlockBytesY = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + StackBuffer xb, yb, result; + check(q5_0_vec_dot(xb.blocks_struct(vx, nb, kBlockBytesX), + yb.blocks_bytes(vy, nb, kBlockBytesY), + result.scalar_f32(s)), + "q5_0_vec_dot"); +} + +// +// Q5_1 -- block size 32, 24 bytes/block (2 delta + 2 min + 4 qh + 16 packed nibbles). +// + +void ggml_quants_halide_quantize_q5_1(const float *x, void *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 4 + 4 + kQK / 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + auto blocks = struct_block_buffer(y, nb, kBlockBytes); + check(q5_1_quantize(xb, blocks), "q5_1_quantize"); +} + +void ggml_quants_halide_dequantize_q5_1(const void *x, float *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 4 + 4 + kQK / 2; + const int32_t nb = static_cast(k / kQK); + auto blocks = struct_block_buffer(x, nb, kBlockBytes); + Buffer yb(y, static_cast(k)); + check(q5_1_dequantize(blocks, yb), "q5_1_dequantize"); +} + +void ggml_quants_halide_vec_dot_q5_1_q8_1(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 32, kBlockBytesX = 4 + 4 + kQK / 2, kBlockBytesY = 4 + kQK; + const int32_t nb = static_cast(n / kQK); + StackBuffer xb, yb, sb, result; + check(q5_1_vec_dot(xb.blocks_struct(vx, nb, kBlockBytesX), + yb.blocks_bytes(vy, nb, kBlockBytesY), + sb.blocks_field_f16(vy, nb, 2, kBlockBytesY), // Q8_1 stored `s` + result.scalar_f32(s)), + "q5_1_vec_dot"); +} + +// +// Q8_0 -- block size 32, 34 bytes/block (2 delta + 32 int8 values). +// + +void ggml_quants_halide_quantize_q8_0(const float *x, void *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 2 + kQK; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + auto blocks = struct_block_buffer(y, nb, kBlockBytes); + check(q8_0_quantize(xb, blocks), "q8_0_quantize"); +} + +void ggml_quants_halide_dequantize_q8_0(const void *x, float *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 2 + kQK; + const int32_t nb = static_cast(k / kQK); + auto blocks = struct_block_buffer(x, nb, kBlockBytes); + Buffer yb(y, static_cast(k)); + check(q8_0_dequantize(blocks, yb), "q8_0_dequantize"); +} + +void ggml_quants_halide_vec_dot_q8_0_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 32, kBlockBytes = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + StackBuffer xb, yb, result; + check(q8_0_vec_dot(xb.blocks_struct(vx, nb, kBlockBytes), // weight: struct-typed + yb.blocks_bytes(vy, nb, kBlockBytes), + result.scalar_f32(s)), + "q8_0_vec_dot"); +} + +// +// Q8_1 -- block size 32, 36 bytes/block (2 delta + 2 sum + 32 int8 values). +// Quantize only -- GGML has no public dequantize for this activation-only format. +// + +void ggml_quants_halide_quantize_q8_1(const float *x, void *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 4 + kQK; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(q8_1_quantize(xb, blocks), "q8_1_quantize"); +} + +// +// Q8_K -- superblock size 256, 292 bytes/block (4 float32 delta + 256 int8 +// values + 16 int16 bsums). Quantize only -- GGML has no public dequantize +// for this activation-only format. +// + +void ggml_quants_halide_quantize_q8_k(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 4 + kQK + (kQK / 16) * 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(q8_k_quantize(xb, blocks), "q8_k_quantize"); +} + +// +// Q2_K -- superblock size 256, 84 bytes/block (16 scale/min nibble bytes + +// 64 packed-2-bit bytes + 2 delta + 2 dmin). Dequantize is native Halide; +// quantize calls out to GGML's own reference (see ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_q2_k(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 16 + kQK / 4 + 4; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(q2_k_quantize(xb, blocks), "q2_k_quantize"); +} + +void ggml_quants_halide_dequantize_q2_k(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 16 + kQK / 4 + 4; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(q2_k_dequantize(blocks, yb), "q2_k_dequantize"); +} + +void ggml_quants_halide_vec_dot_q2_k_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 16 + kQK / 4 + 4, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(q2_k_vec_dot(xb, yb, result), "q2_k_vec_dot"); +} + +// +// Q6_K -- superblock size 256, 210 bytes/block (128 ql + 64 qh + 16 signed +// int8 scales + 2 delta). Dequantize is native Halide; quantize calls out +// to GGML's own reference (see ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_q6_k(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = kQK / 2 + kQK / 4 + kQK / 16 + 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(q6_k_quantize(xb, blocks), "q6_k_quantize"); +} + +void ggml_quants_halide_dequantize_q6_k(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = kQK / 2 + kQK / 4 + kQK / 16 + 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(q6_k_dequantize(blocks, yb), "q6_k_dequantize"); +} + +void ggml_quants_halide_vec_dot_q6_k_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = kQK / 2 + kQK / 4 + kQK / 16 + 2, + kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(q6_k_vec_dot(xb, yb, result), "q6_k_vec_dot"); +} + +// +// Q4_K -- superblock size 256, 144 bytes/block (2 delta + 2 dmin + 12 packed +// scale/min bytes + 128 packed-4-bit bytes). Dequantize is native Halide; +// quantize calls out to GGML's own reference (see ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_q4_k(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 4 + 12 + kQK / 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(q4_k_quantize(xb, blocks), "q4_k_quantize"); +} + +void ggml_quants_halide_dequantize_q4_k(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 4 + 12 + kQK / 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(q4_k_dequantize(blocks, yb), "q4_k_dequantize"); +} + +void ggml_quants_halide_vec_dot_q4_k_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 4 + 12 + kQK / 2, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(q4_k_vec_dot(xb, yb, result), "q4_k_vec_dot"); +} + +// +// Q5_K -- superblock size 256, 176 bytes/block (2 delta + 2 dmin + 12 packed +// scale/min bytes + 32 high-bit bytes + 128 packed-4-bit bytes). Dequantize +// is native Halide; quantize calls out to GGML's own reference (see +// ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_q5_k(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 4 + 12 + kQK / 8 + kQK / 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(q5_k_quantize(xb, blocks), "q5_k_quantize"); +} + +void ggml_quants_halide_dequantize_q5_k(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 4 + 12 + kQK / 8 + kQK / 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(q5_k_dequantize(blocks, yb), "q5_k_dequantize"); +} + +void ggml_quants_halide_vec_dot_q5_k_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 4 + 12 + kQK / 8 + kQK / 2, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(q5_k_vec_dot(xb, yb, result), "q5_k_vec_dot"); +} + +// +// Q3_K -- superblock size 256, 110 bytes/block (32 hmask + 64 packed-2-bit +// bytes + 12 packed scale bytes + 2 delta). Dequantize is native Halide; +// quantize calls out to GGML's own reference (see ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_q3_k(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = kQK / 8 + kQK / 4 + 12 + 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(q3_k_quantize(xb, blocks), "q3_k_quantize"); +} + +void ggml_quants_halide_dequantize_q3_k(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = kQK / 8 + kQK / 4 + 12 + 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(q3_k_dequantize(blocks, yb), "q3_k_dequantize"); +} + +void ggml_quants_halide_vec_dot_q3_k_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = kQK / 8 + kQK / 4 + 12 + 2, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(q3_k_vec_dot(xb, yb, result), "q3_k_vec_dot"); +} + +// +// Q1_0 -- block size 128, 18 bytes/block (2 delta + 16 sign-bit bytes). +// Closed-form both directions, fully native. +// + +void ggml_quants_halide_quantize_q1_0(const float *x, void *y, int64_t k) { + constexpr int kQK = 128, kBlockBytes = 2 + kQK / 8; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + auto blocks = struct_block_buffer(y, nb, kBlockBytes); + check(q1_0_quantize(xb, blocks), "q1_0_quantize"); +} + +void ggml_quants_halide_dequantize_q1_0(const void *x, float *y, int64_t k) { + constexpr int kQK = 128, kBlockBytes = 2 + kQK / 8; + const int32_t nb = static_cast(k / kQK); + auto blocks = struct_block_buffer(x, nb, kBlockBytes); + Buffer yb(y, static_cast(k)); + check(q1_0_dequantize(blocks, yb), "q1_0_dequantize"); +} + +void ggml_quants_halide_vec_dot_q1_0_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + // Q1_0's own block size (128) differs from Q8_0's (32) -- unlike the + // same-granularity types above, nb must be computed separately per side. + constexpr int kQKX = 128, kBlockBytesX = 2 + kQKX / 8; + constexpr int kQKY = 32, kBlockBytesY = 2 + kQKY; + const int32_t nbx = static_cast(n / kQKX); + const int32_t nby = static_cast(n / kQKY); + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nby, kBlockBytesY}}; + auto xb = struct_block_buffer(vx, nbx, kBlockBytesX); // weight: struct-typed + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(q1_0_vec_dot(xb, yb, result), "q1_0_vec_dot"); +} + +// +// MXFP4 -- block size 32, 17 bytes/block (1 E8M0 exponent + 16 packed-4-bit +// codebook-index bytes). Dequantize is native Halide; quantize calls out to +// GGML's own reference (see ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_mxfp4(const float *x, void *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 1 + kQK / 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(mxfp4_quantize(xb, blocks), "mxfp4_quantize"); +} + +void ggml_quants_halide_dequantize_mxfp4(const void *x, float *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 1 + kQK / 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(mxfp4_dequantize(blocks, yb), "mxfp4_dequantize"); +} + +void ggml_quants_halide_vec_dot_mxfp4_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 32, kBlockBytesX = 1 + kQK / 2, kBlockBytesY = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(mxfp4_vec_dot(xb, yb, result), "mxfp4_vec_dot"); +} + +// +// NVFP4 -- block size 64, 36 bytes/block (4 UE4M3 scales + 32 packed-4-bit +// codebook-index bytes). Dequantize is native Halide; quantize calls out to +// GGML's own reference (see ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_nvfp4(const float *x, void *y, int64_t k) { + constexpr int kQK = 64, kSub = 16, kBlockBytes = kQK / kSub + kQK / 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(nvfp4_quantize(xb, blocks), "nvfp4_quantize"); +} + +void ggml_quants_halide_dequantize_nvfp4(const void *x, float *y, int64_t k) { + constexpr int kQK = 64, kSub = 16, kBlockBytes = kQK / kSub + kQK / 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(nvfp4_dequantize(blocks, yb), "nvfp4_dequantize"); +} + +void ggml_quants_halide_vec_dot_nvfp4_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + // NVFP4's own block size (64) differs from Q8_0's (32) -- see Q1_0's + // vec_dot wrapper above for why nb must be computed separately per side. + constexpr int kQKX = 64, kSub = 16, kBlockBytesX = kQKX / kSub + kQKX / 2; + constexpr int kQKY = 32, kBlockBytesY = 2 + kQKY; + const int32_t nbx = static_cast(n / kQKX); + const int32_t nby = static_cast(n / kQKY); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nbx, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nby, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(nvfp4_vec_dot(xb, yb, result), "nvfp4_vec_dot"); +} + +// +// IQ4_NL -- block size 32, 18 bytes/block (2 delta + 16 packed-4-bit +// codebook-index bytes). Dequantize is native Halide; quantize calls out to +// GGML's own reference (see ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_iq4_nl(const float *x, void *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 2 + kQK / 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(iq4_nl_quantize(xb, blocks), "iq4_nl_quantize"); +} + +void ggml_quants_halide_dequantize_iq4_nl(const void *x, float *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 2 + kQK / 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(iq4_nl_dequantize(blocks, yb), "iq4_nl_dequantize"); +} + +void ggml_quants_halide_vec_dot_iq4_nl_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 32, kBlockBytesX = 2 + kQK / 2, kBlockBytesY = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(iq4_nl_vec_dot(xb, yb, result), "iq4_nl_vec_dot"); +} + +// +// IQ4_XS -- superblock size 256, 136 bytes/block (2 delta + 2 scales_h + 4 +// scales_l + 128 packed-4-bit codebook-index bytes). Dequantize is native +// Halide; quantize calls out to GGML's own reference (see +// ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_iq4_xs(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + 2 + 4 + kQK / 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(iq4_xs_quantize(xb, blocks), "iq4_xs_quantize"); +} + +void ggml_quants_halide_dequantize_iq4_xs(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + 2 + 4 + kQK / 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(iq4_xs_dequantize(blocks, yb), "iq4_xs_dequantize"); +} + +void ggml_quants_halide_vec_dot_iq4_xs_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 2 + 2 + 4 + kQK / 2, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(iq4_xs_vec_dot(xb, yb, result), "iq4_xs_vec_dot"); +} + +// +// TQ1_0 -- superblock size 256, 54 bytes/block (48 base-3-packed qs + 4 +// base-3-packed qh + 2 delta). Dequantize is native Halide; quantize calls +// out to GGML's own reference (see ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_tq1_0(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 54; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(tq1_0_quantize(xb, blocks), "tq1_0_quantize"); +} + +void ggml_quants_halide_dequantize_tq1_0(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 54; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(tq1_0_dequantize(blocks, yb), "tq1_0_dequantize"); +} + +void ggml_quants_halide_vec_dot_tq1_0_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 54, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(tq1_0_vec_dot(xb, yb, result), "tq1_0_vec_dot"); +} + +// +// TQ2_0 -- superblock size 256, 66 bytes/block (64 packed-2-bit qs + 2 +// delta -- qs before d, unlike every other type here). Dequantize is native +// Halide; quantize calls out to GGML's own reference (see +// ggml_extern_quantize.cpp). +// + +void ggml_quants_halide_quantize_tq2_0(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = kQK / 4 + 2; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(tq2_0_quantize(xb, blocks), "tq2_0_quantize"); +} + +void ggml_quants_halide_dequantize_tq2_0(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = kQK / 4 + 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(tq2_0_dequantize(blocks, yb), "tq2_0_dequantize"); +} + +void ggml_quants_halide_vec_dot_tq2_0_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = kQK / 4 + 2, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(tq2_0_vec_dot(xb, yb, result), "tq2_0_vec_dot"); +} + +// +// IQ2_XXS -- superblock size 256, 66 bytes/block (2 delta + 64 packed qs). +// Dequantize only -- see ggml_quants.h for why there's no quantize here. +// + +void ggml_quants_halide_dequantize_iq2_xxs(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + kQK / 8 * 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(iq2_xxs_dequantize(blocks, yb), "iq2_xxs_dequantize"); +} + +void ggml_quants_halide_vec_dot_iq2_xxs_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 2 + kQK / 8 * 2, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(iq2_xxs_vec_dot(xb, yb, result), "iq2_xxs_vec_dot"); +} + +// +// IQ2_XS -- superblock size 256, 74 bytes/block. Dequantize only. +// + +void ggml_quants_halide_dequantize_iq2_xs(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + kQK / 8 * 2 + kQK / 32; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(iq2_xs_dequantize(blocks, yb), "iq2_xs_dequantize"); +} + +void ggml_quants_halide_vec_dot_iq2_xs_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 2 + kQK / 8 * 2 + kQK / 32, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(iq2_xs_vec_dot(xb, yb, result), "iq2_xs_vec_dot"); +} + +// +// IQ2_S -- superblock size 256, 82 bytes/block. Dequantize is native +// Halide; quantize calls out to GGML's own reference. +// + +void ggml_quants_halide_quantize_iq2_s(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + kQK / 8 + kQK / 8 + kQK / 32 + kQK / 32; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(iq2_s_quantize(xb, blocks), "iq2_s_quantize"); +} + +void ggml_quants_halide_dequantize_iq2_s(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + kQK / 8 + kQK / 8 + kQK / 32 + kQK / 32; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(iq2_s_dequantize(blocks, yb), "iq2_s_dequantize"); +} + +void ggml_quants_halide_vec_dot_iq2_s_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 2 + kQK / 8 + kQK / 8 + kQK / 32 + kQK / 32, + kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(iq2_s_vec_dot(xb, yb, result), "iq2_s_vec_dot"); +} + +// +// IQ3_XXS -- superblock size 256, 98 bytes/block. Dequantize is native +// Halide; quantize calls out to GGML's own reference. +// + +void ggml_quants_halide_quantize_iq3_xxs(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + 3 * kQK / 8; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(iq3_xxs_quantize(xb, blocks), "iq3_xxs_quantize"); +} + +void ggml_quants_halide_dequantize_iq3_xxs(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + 3 * kQK / 8; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(iq3_xxs_dequantize(blocks, yb), "iq3_xxs_dequantize"); +} + +void ggml_quants_halide_vec_dot_iq3_xxs_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 2 + 3 * kQK / 8, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(iq3_xxs_vec_dot(xb, yb, result), "iq3_xxs_vec_dot"); +} + +// +// IQ3_S -- superblock size 256, 110 bytes/block. Dequantize is native +// Halide; quantize calls out to GGML's own reference. +// + +void ggml_quants_halide_quantize_iq3_s(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + kQK / 4 + kQK / 32 + kQK / 8 + kQK / 64; + Buffer xb(const_cast(x), static_cast(k)); + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, shape); + check(iq3_s_quantize(xb, blocks), "iq3_s_quantize"); +} + +void ggml_quants_halide_dequantize_iq3_s(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + kQK / 4 + kQK / 32 + kQK / 8 + kQK / 64; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(iq3_s_dequantize(blocks, yb), "iq3_s_dequantize"); +} + +void ggml_quants_halide_vec_dot_iq3_s_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 2 + kQK / 4 + kQK / 32 + kQK / 8 + kQK / 64, + kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(iq3_s_vec_dot(xb, yb, result), "iq3_s_vec_dot"); +} + +// +// IQ1_S -- superblock size 256, 50 bytes/block. Dequantize only. +// + +void ggml_quants_halide_dequantize_iq1_s(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 2 + kQK / 8 + kQK / 16; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(iq1_s_dequantize(blocks, yb), "iq1_s_dequantize"); +} + +void ggml_quants_halide_vec_dot_iq1_s_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = 2 + kQK / 8 + kQK / 16, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(iq1_s_vec_dot(xb, yb, result), "iq1_s_vec_dot"); +} + +// +// IQ1_M -- superblock size 256, 56 bytes/block. Dequantize only. +// + +void ggml_quants_halide_dequantize_iq1_m(const void *x, float *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = kQK / 8 + kQK / 16 + kQK / 32; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t shape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(const_cast(static_cast(x)), 2, shape); + Buffer yb(y, static_cast(k)); + check(iq1_m_dequantize(blocks, yb), "iq1_m_dequantize"); +} + +void ggml_quants_halide_vec_dot_iq1_m_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc) { + constexpr int kQK = 256, kBlockBytesX = kQK / 8 + kQK / 16 + kQK / 32, kBlockBytesY = 4 + kQK + (kQK / 16) * 2; + const int32_t nb = static_cast(n / kQK); + halide_dimension_t xshape[2] = {{0, kBlockBytesX, 1}, {0, nb, kBlockBytesX}}; + halide_dimension_t yshape[2] = {{0, kBlockBytesY, 1}, {0, nb, kBlockBytesY}}; + Buffer xb(const_cast(static_cast(vx)), 2, xshape); + Buffer yb(const_cast(static_cast(vy)), 2, yshape); + Buffer result = Buffer::make_scalar(s); + check(iq1_m_vec_dot(xb, yb, result), "iq1_m_vec_dot"); +} + +// +// F16 -- block size 1, 2 bytes/element (plain IEEE binary16 cast, no header). +// + +void ggml_quants_halide_quantize_f16(const float *x, void *y, int64_t k) { + Buffer xb(const_cast(x), static_cast(k)); + Buffer yb(static_cast(y), static_cast(k)); + check(f16_quantize(xb, yb), "f16_quantize"); +} + +void ggml_quants_halide_dequantize_f16(const void *x, float *y, int64_t k) { + Buffer xb(const_cast(static_cast(x)), static_cast(k)); + Buffer yb(y, static_cast(k)); + check(f16_dequantize(xb, yb), "f16_dequantize"); +} + +// +// BF16 -- block size 1, 2 bytes/element (plain bfloat16 cast, no header). +// + +void ggml_quants_halide_quantize_bf16(const float *x, void *y, int64_t k) { + Buffer xb(const_cast(x), static_cast(k)); + Buffer yb(static_cast(y), static_cast(k)); + check(bf16_quantize(xb, yb), "bf16_quantize"); +} + +void ggml_quants_halide_dequantize_bf16(const void *x, float *y, int64_t k) { + Buffer xb(const_cast(static_cast(x)), static_cast(k)); + Buffer yb(y, static_cast(k)); + check(bf16_dequantize(xb, yb), "bf16_dequantize"); +} + +// +// Repack quantize_mat -- interleaves 4 contiguous rows of `k` floats (row r +// at x[r*k .. r*k+k)) into one packed activation-format block per chunk. +// `x` is wrapped as a 2-D buffer (dim 0: column-within-row, extent k; dim 1: +// row, extent 4, stride k) so the generator can address x_(col, row) +// directly instead of hand-computing `row*k + col`. +// + +void ggml_quants_halide_repack_quantize_mat_q8_0_4x4(const float *x, void *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 4 * 2 + kQK * 4; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t xshape[2] = {{0, static_cast(k), 1}, {0, 4, static_cast(k)}}; + Buffer xb(const_cast(x), 2, xshape); + halide_dimension_t yshape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, yshape); + check(q8_0_4x4_quantize_mat(xb, blocks), "q8_0_4x4_quantize_mat"); +} + +void ggml_quants_halide_repack_quantize_mat_q8_0_4x8(const float *x, void *y, int64_t k) { + constexpr int kQK = 32, kBlockBytes = 4 * 2 + kQK * 4; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t xshape[2] = {{0, static_cast(k), 1}, {0, 4, static_cast(k)}}; + Buffer xb(const_cast(x), 2, xshape); + halide_dimension_t yshape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, yshape); + check(q8_0_4x8_quantize_mat(xb, blocks), "q8_0_4x8_quantize_mat"); +} + +void ggml_quants_halide_repack_quantize_mat_q8_k_4x4(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 4 * 4 + kQK * 4 + (kQK / 16) * 4 * 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t xshape[2] = {{0, static_cast(k), 1}, {0, 4, static_cast(k)}}; + Buffer xb(const_cast(x), 2, xshape); + halide_dimension_t yshape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, yshape); + check(q8_k_4x4_quantize_mat(xb, blocks), "q8_k_4x4_quantize_mat"); +} + +void ggml_quants_halide_repack_quantize_mat_q8_k_4x8(const float *x, void *y, int64_t k) { + constexpr int kQK = 256, kBlockBytes = 4 * 4 + kQK * 4 + (kQK / 16) * 4 * 2; + const int32_t nb = static_cast(k / kQK); + halide_dimension_t xshape[2] = {{0, static_cast(k), 1}, {0, 4, static_cast(k)}}; + Buffer xb(const_cast(x), 2, xshape); + halide_dimension_t yshape[2] = {{0, kBlockBytes, 1}, {0, nb, kBlockBytes}}; + Buffer blocks(static_cast(y), 2, yshape); + check(q8_k_4x8_quantize_mat(xb, blocks), "q8_k_4x8_quantize_mat"); +} + +// +// Repack gemv/gemm: dot a repack-interleaved weight matrix against a Q8_0 +// activation row (gemv) or `nr` rows packed 4 at a time (gemm). See +// repack_gemv_generators.cpp/repack_gemm_generators.cpp for the packed +// buffer layouts these shapes describe; `bs` (gemm's output row stride) is +// unused by gemv, matching gemx_fn_t's own nr == 1 convention. +// + +void ggml_quants_halide_repack_gemv_q4_0_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q4_0_4x4_gemv(wb, ab, sb), "q4_0_4x4_gemv"); +} + +void ggml_quants_halide_repack_gemv_q4_0_4x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q4_0_4x8_gemv(wb, ab, sb), "q4_0_4x8_gemv"); +} + +void ggml_quants_halide_repack_gemv_q4_0_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 32, kNCols = 8, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q4_0_8x8_gemv(wb, ab, sb), "q4_0_8x8_gemv"); +} + +void ggml_quants_halide_repack_gemv_q8_0_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 8) / 8; + constexpr int kActBlockBytes = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q8_0_4x4_gemv(wb, ab, sb), "q8_0_4x4_gemv"); +} + +void ggml_quants_halide_repack_gemv_q8_0_4x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 8) / 8; + constexpr int kActBlockBytes = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q8_0_4x8_gemv(wb, ab, sb), "q8_0_4x8_gemv"); +} + +void ggml_quants_halide_repack_gemm_q4_0_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActNRows = 4, kActBlockBytes = 2 * kActNRows + kQK * kActNRows; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = { + {0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}, {0, nr_groups, kActBlockBytes * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q4_0_4x4_gemm(wb, ab, sb), "q4_0_4x4_gemm"); +} + +void ggml_quants_halide_repack_gemm_q4_0_4x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActNRows = 4, kActBlockBytes = 2 * kActNRows + kQK * kActNRows; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = { + {0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}, {0, nr_groups, kActBlockBytes * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q4_0_4x8_gemm(wb, ab, sb), "q4_0_4x8_gemm"); +} + +void ggml_quants_halide_repack_gemm_q4_0_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 32, kNCols = 8, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActNRows = 4, kActBlockBytes = 2 * kActNRows + kQK * kActNRows; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = { + {0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}, {0, nr_groups, kActBlockBytes * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q4_0_8x8_gemm(wb, ab, sb), "q4_0_8x8_gemm"); +} + +void ggml_quants_halide_repack_gemm_q8_0_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 8) / 8; + constexpr int kActNRows = 4, kActBlockBytes = 2 * kActNRows + kQK * kActNRows; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = { + {0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}, {0, nr_groups, kActBlockBytes * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q8_0_4x4_gemm(wb, ab, sb), "q8_0_4x4_gemm"); +} + +void ggml_quants_halide_repack_gemm_q8_0_4x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 8) / 8; + constexpr int kActNRows = 4, kActBlockBytes = 2 * kActNRows + kQK * kActNRows; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = { + {0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}, {0, nr_groups, kActBlockBytes * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q8_0_4x8_gemm(wb, ab, sb), "q8_0_4x8_gemm"); +} + +// +// IQ4_NL/MXFP4 repack gemv/gemm: same 3-D packed-buffer shapes as Q4_0's +// above, but MXFP4's weight header is N bytes (1 E8M0 exponent per column) +// instead of 2*N (an fp16 delta per column) -- see repack_gemv_generators.cpp. +// + +void ggml_quants_halide_repack_gemv_iq4_nl_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(iq4_nl_4x4_gemv(wb, ab, sb), "iq4_nl_4x4_gemv"); +} + +void ggml_quants_halide_repack_gemv_iq4_nl_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 32, kNCols = 8, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(iq4_nl_8x8_gemv(wb, ab, sb), "iq4_nl_8x8_gemv"); +} + +void ggml_quants_halide_repack_gemv_mxfp4_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(mxfp4_4x4_gemv(wb, ab, sb), "mxfp4_4x4_gemv"); +} + +void ggml_quants_halide_repack_gemv_mxfp4_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 32, kNCols = 8, kWeightBlockBytes = kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 2 + kQK; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(mxfp4_8x8_gemv(wb, ab, sb), "mxfp4_8x8_gemv"); +} + +void ggml_quants_halide_repack_gemm_iq4_nl_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActNRows = 4, kActBlockBytes = 2 * kActNRows + kQK * kActNRows; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = { + {0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}, {0, nr_groups, kActBlockBytes * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(iq4_nl_4x4_gemm(wb, ab, sb), "iq4_nl_4x4_gemm"); +} + +void ggml_quants_halide_repack_gemm_iq4_nl_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 32, kNCols = 8, kWeightBlockBytes = 2 * kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActNRows = 4, kActBlockBytes = 2 * kActNRows + kQK * kActNRows; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = { + {0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}, {0, nr_groups, kActBlockBytes * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(iq4_nl_8x8_gemm(wb, ab, sb), "iq4_nl_8x8_gemm"); +} + +void ggml_quants_halide_repack_gemm_mxfp4_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 32, kNCols = 4, kWeightBlockBytes = kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActNRows = 4, kActBlockBytes = 2 * kActNRows + kQK * kActNRows; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = { + {0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}, {0, nr_groups, kActBlockBytes * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(mxfp4_4x4_gemm(wb, ab, sb), "mxfp4_4x4_gemm"); +} + +void ggml_quants_halide_repack_gemm_mxfp4_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 32, kNCols = 8, kWeightBlockBytes = kNCols + (kQK * kNCols * 4) / 8; + constexpr int kActNRows = 4, kActBlockBytes = 2 * kActNRows + kQK * kActNRows; + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = { + {0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}, {0, nr_groups, kActBlockBytes * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(mxfp4_8x8_gemm(wb, ab, sb), "mxfp4_8x8_gemm"); +} + +// +// Q4_K repack gemv/gemm: 256-element superblocks, always 8 interleaved +// columns, paired with Q8_K activations. See repack_gemv_generators.cpp/ +// repack_gemm_generators.cpp for the packed buffer layouts. +// + +void ggml_quants_halide_repack_gemv_q4_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 256, kNCols = 8, kWeightBlockBytes = 16 + 16 + 96 + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 4 + kQK + 2 * (kQK / 16); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q4_k_8x4_gemv(wb, ab, sb), "q4_k_8x4_gemv"); +} + +void ggml_quants_halide_repack_gemv_q4_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 256, kNCols = 8, kWeightBlockBytes = 16 + 16 + 96 + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 4 + kQK + 2 * (kQK / 16); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q4_k_8x8_gemv(wb, ab, sb), "q4_k_8x8_gemv"); +} + +void ggml_quants_halide_repack_gemm_q4_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 256, kNCols = 8, kWeightBlockBytes = 16 + 16 + 96 + (kQK * kNCols * 4) / 8; + // block_q8_Kx4's real per-block size is 1168 bytes (16 header + 1024 qs + + // 128 bsums) -- that's the actual stride between consecutive blocks in + // memory, even though the Halide kernel only ever reads the first 1040 + // bytes (header + qs; bsums are unused, see repack_gemm_generators.cpp). + constexpr int kActNRows = 4, kActBlockBytesUsed = 4 * kActNRows + kQK * kActNRows; + constexpr int kActBlockBytesReal = kActBlockBytesUsed + 2 * (kQK / 4); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = {{0, kActBlockBytesUsed, 1}, + {0, nb, kActBlockBytesReal}, + {0, nr_groups, kActBlockBytesReal * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q4_k_8x4_gemm(wb, ab, sb), "q4_k_8x4_gemm"); +} + +void ggml_quants_halide_repack_gemm_q4_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 256, kNCols = 8, kWeightBlockBytes = 16 + 16 + 96 + (kQK * kNCols * 4) / 8; + constexpr int kActNRows = 4, kActBlockBytesUsed = 4 * kActNRows + kQK * kActNRows; + constexpr int kActBlockBytesReal = kActBlockBytesUsed + 2 * (kQK / 4); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = {{0, kActBlockBytesUsed, 1}, + {0, nb, kActBlockBytesReal}, + {0, nr_groups, kActBlockBytesReal * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q4_k_8x8_gemm(wb, ab, sb), "q4_k_8x8_gemm"); +} + +// +// Q5_K repack gemv/gemm: same 256-element superblock/8-column structure as +// Q4_K, plus a 256-byte qh (5th bit) array between scales and qs. +// + +void ggml_quants_halide_repack_gemv_q5_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 256, kNCols = 8; + constexpr int kWeightBlockBytes = 16 + 16 + 96 + (kQK * kNCols) / 8 + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 4 + kQK + 2 * (kQK / 16); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q5_k_8x4_gemv(wb, ab, sb), "q5_k_8x4_gemv"); +} + +void ggml_quants_halide_repack_gemv_q5_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 256, kNCols = 8; + constexpr int kWeightBlockBytes = 16 + 16 + 96 + (kQK * kNCols) / 8 + (kQK * kNCols * 4) / 8; + constexpr int kActBlockBytes = 4 + kQK + 2 * (kQK / 16); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q5_k_8x8_gemv(wb, ab, sb), "q5_k_8x8_gemv"); +} + +void ggml_quants_halide_repack_gemm_q5_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 256, kNCols = 8; + constexpr int kWeightBlockBytes = 16 + 16 + 96 + (kQK * kNCols) / 8 + (kQK * kNCols * 4) / 8; + constexpr int kActNRows = 4, kActBlockBytesUsed = 4 * kActNRows + kQK * kActNRows; + constexpr int kActBlockBytesReal = kActBlockBytesUsed + 2 * (kQK / 4); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = {{0, kActBlockBytesUsed, 1}, + {0, nb, kActBlockBytesReal}, + {0, nr_groups, kActBlockBytesReal * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q5_k_8x4_gemm(wb, ab, sb), "q5_k_8x4_gemm"); +} + +void ggml_quants_halide_repack_gemm_q5_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 256, kNCols = 8; + constexpr int kWeightBlockBytes = 16 + 16 + 96 + (kQK * kNCols) / 8 + (kQK * kNCols * 4) / 8; + constexpr int kActNRows = 4, kActBlockBytesUsed = 4 * kActNRows + kQK * kActNRows; + constexpr int kActBlockBytesReal = kActBlockBytesUsed + 2 * (kQK / 4); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = {{0, kActBlockBytesUsed, 1}, + {0, nb, kActBlockBytesReal}, + {0, nr_groups, kActBlockBytesReal * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q5_k_8x8_gemm(wb, ab, sb), "q5_k_8x8_gemm"); +} + +// +// Q6_K repack gemv/gemm: plain signed-int8-per-sub-group scales, no compact +// bit packing -- see repack_gemv_generators.cpp/repack_gemm_generators.cpp. +// + +void ggml_quants_halide_repack_gemv_q6_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 256, kNCols = 8; + constexpr int kWeightBlockBytes = 16 + 128 + (kQK * kNCols * 4) / 8 + (kQK * kNCols * 2) / 8; + constexpr int kActBlockBytes = 4 + kQK + 2 * (kQK / 16); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q6_k_8x4_gemv(wb, ab, sb), "q6_k_8x4_gemv"); +} + +void ggml_quants_halide_repack_gemv_q6_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 256, kNCols = 8; + constexpr int kWeightBlockBytes = 16 + 128 + (kQK * kNCols * 4) / 8 + (kQK * kNCols * 2) / 8; + constexpr int kActBlockBytes = 4 + kQK + 2 * (kQK / 16); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q6_k_8x8_gemv(wb, ab, sb), "q6_k_8x8_gemv"); +} + +void ggml_quants_halide_repack_gemm_q6_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 256, kNCols = 8; + constexpr int kWeightBlockBytes = 16 + 128 + (kQK * kNCols * 4) / 8 + (kQK * kNCols * 2) / 8; + constexpr int kActNRows = 4, kActBlockBytesUsed = 4 * kActNRows + kQK * kActNRows; + constexpr int kActBlockBytesReal = kActBlockBytesUsed + 2 * (kQK / 4); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = {{0, kActBlockBytesUsed, 1}, + {0, nb, kActBlockBytesReal}, + {0, nr_groups, kActBlockBytesReal * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q6_k_8x4_gemm(wb, ab, sb), "q6_k_8x4_gemm"); +} + +void ggml_quants_halide_repack_gemm_q6_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 256, kNCols = 8; + constexpr int kWeightBlockBytes = 16 + 128 + (kQK * kNCols * 4) / 8 + (kQK * kNCols * 2) / 8; + constexpr int kActNRows = 4, kActBlockBytesUsed = 4 * kActNRows + kQK * kActNRows; + constexpr int kActBlockBytesReal = kActBlockBytesUsed + 2 * (kQK / 4); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = {{0, kActBlockBytesUsed, 1}, + {0, nb, kActBlockBytesReal}, + {0, nr_groups, kActBlockBytesReal * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q6_k_8x8_gemm(wb, ab, sb), "q6_k_8x8_gemm"); +} + +// +// Q2_K repack gemv/gemm: only one registered variant (8x8). See +// repack_gemv_generators.cpp/repack_gemm_generators.cpp. +// + +void ggml_quants_halide_repack_gemv_q2_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + (void)bs; + (void)nr; + constexpr int kQK = 256, kNCols = 8; + constexpr int kWeightBlockBytes = 16 + 16 + 128 + (kQK * kNCols * 2) / 8; + constexpr int kActBlockBytes = 4 + kQK + 2 * (kQK / 16); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[2] = {{0, kActBlockBytes, 1}, {0, nb, kActBlockBytes}}; + Buffer ab(const_cast(static_cast(vy)), 2, ashape); + halide_dimension_t sshape[2] = {{0, kNCols, 1}, {0, nc_groups, kNCols}}; + Buffer sb(s, 2, sshape); + check(q2_k_8x8_gemv(wb, ab, sb), "q2_k_8x8_gemv"); +} + +void ggml_quants_halide_repack_gemm_q2_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc) { + constexpr int kQK = 256, kNCols = 8; + constexpr int kWeightBlockBytes = 16 + 16 + 128 + (kQK * kNCols * 2) / 8; + constexpr int kActNRows = 4, kActBlockBytesUsed = 4 * kActNRows + kQK * kActNRows; + constexpr int kActBlockBytesReal = kActBlockBytesUsed + 2 * (kQK / 4); + const int32_t nb = static_cast(n / kQK); + const int32_t nc_groups = static_cast(nc / kNCols); + const int32_t nr_groups = static_cast(nr / kActNRows); + const int32_t bs32 = static_cast(bs); + halide_dimension_t wshape[3] = { + {0, kWeightBlockBytes, 1}, {0, nb, kWeightBlockBytes}, {0, nc_groups, kWeightBlockBytes * nb}}; + Buffer wb(const_cast(static_cast(vx)), 3, wshape); + halide_dimension_t ashape[3] = {{0, kActBlockBytesUsed, 1}, + {0, nb, kActBlockBytesReal}, + {0, nr_groups, kActBlockBytesReal * nb}}; + Buffer ab(const_cast(static_cast(vy)), 3, ashape); + halide_dimension_t sshape[4] = { + {0, kNCols, 1}, {0, nc_groups, kNCols}, {0, kActNRows, bs32}, {0, nr_groups, kActNRows * bs32}}; + Buffer sb(s, 4, sshape); + check(q2_k_8x8_gemm(wb, ab, sb), "q2_k_8x8_gemm"); +} + +} // extern "C" diff --git a/apps/ggml/halide/ggml_quants.h b/apps/ggml/halide/ggml_quants.h new file mode 100644 index 000000000000..65dce00da119 --- /dev/null +++ b/apps/ggml/halide/ggml_quants.h @@ -0,0 +1,276 @@ +#pragma once + +// Plain C ABI for the Halide-generated quantize/dequantize/vec_dot kernels, +// matching apps/ggml/include/kernel_registry.h's quantize_fn_t/ +// dequantize_fn_t/vec_dot_fn_t signatures exactly (by signature +// compatibility alone -- no shared header needed between this library and +// the benchmark harness). +// +// Q8_1 has no dequantize entry: it's an activation-only format (GGML itself +// has no public to_float for it), so there is nothing to implement. +// +// Every vec_dot__ function computes the dot product between a row of +// weight-type x and a row of activation-type y (y is x's GGML vec_dot_type). +// Both operands flow through the Approximation framework: the generic +// symmetric/lookup_table/k_quant vec_dot generators splice weight and +// activation codecs from quant_components.h via approximate_by/sever +// (see vec_dot_generator_base.h). + +#include +#include + +extern "C" { + +void ggml_quants_halide_quantize_q4_0(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q4_0(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q4_0_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_q4_1(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q4_1(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q4_1_q8_1(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_q5_0(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q5_0(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q5_0_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_q5_1(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q5_1(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q5_1_q8_1(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_q8_0(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q8_0(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q8_0_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_q8_1(const float *x, void *y, int64_t k); + +void ggml_quants_halide_quantize_q8_k(const float *x, void *y, int64_t k); + +// Q2_K, Q6_K: dequantize is a from-scratch Halide implementation; quantize +// is scaffolding that calls out to GGML's own reference (see +// ggml_extern_quantize.cpp) pending a from-scratch port of GGML's iterative +// scale search. vec_dot is from-scratch (against Q8_K activations). +void ggml_quants_halide_quantize_q2_k(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q2_k(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q2_k_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_q6_k(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q6_k(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q6_k_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_q4_k(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q4_k(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q4_k_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_q5_k(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q5_k(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q5_k_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_q3_k(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q3_k(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q3_k_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +// Q1_0: closed-form both directions, fully native. +void ggml_quants_halide_quantize_q1_0(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_q1_0(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_q1_0_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +// MXFP4, NVFP4, IQ4_NL, IQ4_XS, TQ1_0, TQ2_0: dequantize is native Halide; +// quantize calls out to GGML's own reference (see ggml_extern_quantize.cpp). +void ggml_quants_halide_quantize_mxfp4(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_mxfp4(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_mxfp4_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_nvfp4(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_nvfp4(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_nvfp4_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_iq4_nl(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_iq4_nl(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_iq4_nl_q8_0(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_iq4_xs(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_iq4_xs(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_iq4_xs_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_tq1_0(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_tq1_0(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_tq1_0_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_tq2_0(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_tq2_0(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_tq2_0_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +// IQ2_XXS: dequantize only. GGML has no public from_float_ref for this +// importance-matrix-only codebook type (only a private whole-matrix +// quantizer -- see providers/ggml_internal_abi.h), so there is no +// from-scratch quantizer to write against it. vec_dot is still implemented +// (GGML's own reference quantizer is used to produce test/benchmark input). +void ggml_quants_halide_dequantize_iq2_xxs(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_iq2_xxs_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +// IQ2_XS: dequantize only (same reason as IQ2_XXS above). +void ggml_quants_halide_dequantize_iq2_xs(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_iq2_xs_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +// IQ2_S, IQ3_XXS, IQ3_S: dequantize is native Halide; quantize calls out to +// GGML's own reference (see ggml_extern_quantize.cpp) -- these three do +// have a public from_float_ref, unlike IQ2_XXS/IQ2_XS/IQ1_S/IQ1_M. +void ggml_quants_halide_quantize_iq2_s(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_iq2_s(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_iq2_s_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_iq3_xxs(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_iq3_xxs(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_iq3_xxs_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_quantize_iq3_s(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_iq3_s(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_iq3_s_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +// IQ1_S, IQ1_M: dequantize only (same reason as IQ2_XXS above). +void ggml_quants_halide_dequantize_iq1_s(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_iq1_s_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +void ggml_quants_halide_dequantize_iq1_m(const void *x, float *y, int64_t k); +void ggml_quants_halide_vec_dot_iq1_m_q8_k(int n, float *s, size_t bs, const void *vx, size_t bx, const void *vy, + size_t by, int nrc); + +// F16, BF16: not really "quantized" types -- block size 1, a plain per- +// element float cast. Both directions are fully native Halide (Halide's +// built-in float16_t/bfloat16_t casts implement the same IEEE round-to- +// nearest-even conversions GGML's own reference uses). No vec_dot: not part +// of the quantized-format vec_dot sweep this directory otherwise covers. +void ggml_quants_halide_quantize_f16(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_f16(const void *x, float *y, int64_t k); + +void ggml_quants_halide_quantize_bf16(const float *x, void *y, int64_t k); +void ggml_quants_halide_dequantize_bf16(const void *x, float *y, int64_t k); + +// Repack quantize_mat kernels: interleave 4 contiguous rows of `k` floats +// into one packed activation-format block per `k`-sized chunk (see +// repack_quantize_mat_generators.cpp). Signature-compatible with +// quantize_fn_t (same as every other quantize_* above) -- these are just +// keyed by activation format + interleave width rather than weight type, +// since GGML itself only has 4 distinct quantize_mat implementations shared +// across every repack weight type (see k_repack_entries in +// providers/ggml_provider.cpp and this library's registration in +// providers/halide_provider.cpp). +void ggml_quants_halide_repack_quantize_mat_q8_0_4x4(const float *x, void *y, int64_t k); +void ggml_quants_halide_repack_quantize_mat_q8_0_4x8(const float *x, void *y, int64_t k); +void ggml_quants_halide_repack_quantize_mat_q8_k_4x4(const float *x, void *y, int64_t k); +void ggml_quants_halide_repack_quantize_mat_q8_k_4x8(const float *x, void *y, int64_t k); + +// Repack gemv/gemm: dot a repack-interleaved weight matrix (see +// repack_quantize_mat_generators.cpp/repack_gemv_generators.cpp for the +// packed layout) against, respectively, one plain-Q8_0 activation row +// (gemv, matching gemx_fn_t's nr == 1 case) or `nr` activation rows packed 4 +// at a time by the matching repack_quantize_mat_* kernel above (gemm). +// Q4_0's own two variants (4x4, 4x8) share one packed-weight byte layout +// (see repack_gemv_generators.cpp); 8x8 uses a wider one. All 3 use +// activations quantized by the matching-blocklen quantize_mat kernel, same +// as ggml_provider.cpp's own k_repack_entries table. +void ggml_quants_halide_repack_gemv_q4_0_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemv_q4_0_4x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemv_q4_0_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemv_q8_0_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemv_q8_0_4x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); + +void ggml_quants_halide_repack_gemm_q4_0_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q4_0_4x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q4_0_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q8_0_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q8_0_4x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); + +// IQ4_NL/MXFP4: same repack-interleave scheme as Q4_0 (nibble/halves split), +// but a plain codebook lookup (no XOR trick) for the weight value, and only +// two interleave widths each (4x4, 8x8 -- no 4x8), matching GGML's own +// k_repack_entries. +void ggml_quants_halide_repack_gemv_iq4_nl_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc); +void ggml_quants_halide_repack_gemv_iq4_nl_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc); +void ggml_quants_halide_repack_gemv_mxfp4_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc); +void ggml_quants_halide_repack_gemv_mxfp4_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc); + +void ggml_quants_halide_repack_gemm_iq4_nl_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc); +void ggml_quants_halide_repack_gemm_iq4_nl_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc); +void ggml_quants_halide_repack_gemm_mxfp4_4x4_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc); +void ggml_quants_halide_repack_gemm_mxfp4_8x8_q8_0(int n, float *s, size_t bs, const void *vx, const void *vy, + int nr, int nc); + +// Q4_K: 256-element superblocks, always 8 interleaved columns, paired with +// Q8_K activations. +void ggml_quants_halide_repack_gemv_q4_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemv_q4_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q4_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q4_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); + +// Q5_K: same 256-element superblock/8-column structure as Q4_K, plus a 5th +// (high) bit array. +void ggml_quants_halide_repack_gemv_q5_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemv_q5_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q5_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q5_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); + +// Q6_K: plain signed-int8-per-sub-group scales (no compact bit packing). +void ggml_quants_halide_repack_gemv_q6_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemv_q6_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q6_k_8x4_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q6_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); + +// Q2_K: only one registered variant (8x8 -- GGML has no ARM 8x4 path). +void ggml_quants_halide_repack_gemv_q2_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +void ggml_quants_halide_repack_gemm_q2_k_8x8_q8_k(int n, float *s, size_t bs, const void *vx, const void *vy, int nr, + int nc); +} diff --git a/apps/ggml/halide/iq_grids_data.h b/apps/ggml/halide/iq_grids_data.h new file mode 100644 index 000000000000..9076fe4ec162 --- /dev/null +++ b/apps/ggml/halide/iq_grids_data.h @@ -0,0 +1,4770 @@ +// Constant lookup tables copied verbatim from GGML's src/ggml-common.h +// (commit eced84c86f8b012c752c016f7fe789adea168e1e / v0.15.3), used by the +// codebook-based IQ2/IQ3/IQ1 dequantize generators. These are fixed, +// published codebook constants, not derived logic -- transcribed exactly, +// not "included" from GGML (no GGML headers are used elsewhere in this +// directory's generators). +#pragma once +#include +namespace iq_grids { +constexpr uint8_t kmask_iq2xs[8] = { + 1, + 2, + 4, + 8, + 16, + 32, + 64, + 128, +}; +constexpr uint8_t ksigns_iq2xs[128] = { + 0, + 129, + 130, + 3, + 132, + 5, + 6, + 135, + 136, + 9, + 10, + 139, + 12, + 141, + 142, + 15, + 144, + 17, + 18, + 147, + 20, + 149, + 150, + 23, + 24, + 153, + 154, + 27, + 156, + 29, + 30, + 159, + 160, + 33, + 34, + 163, + 36, + 165, + 166, + 39, + 40, + 169, + 170, + 43, + 172, + 45, + 46, + 175, + 48, + 177, + 178, + 51, + 180, + 53, + 54, + 183, + 184, + 57, + 58, + 187, + 60, + 189, + 190, + 63, + 192, + 65, + 66, + 195, + 68, + 197, + 198, + 71, + 72, + 201, + 202, + 75, + 204, + 77, + 78, + 207, + 80, + 209, + 210, + 83, + 212, + 85, + 86, + 215, + 216, + 89, + 90, + 219, + 92, + 221, + 222, + 95, + 96, + 225, + 226, + 99, + 228, + 101, + 102, + 231, + 232, + 105, + 106, + 235, + 108, + 237, + 238, + 111, + 240, + 113, + 114, + 243, + 116, + 245, + 246, + 119, + 120, + 249, + 250, + 123, + 252, + 125, + 126, + 255, +}; +constexpr uint64_t iq2xxs_grid[256] = { + 0x0808080808080808, + 0x080808080808082b, + 0x0808080808081919, + 0x0808080808082b08, + 0x0808080808082b2b, + 0x0808080808190819, + 0x0808080808191908, + 0x08080808082b0808, + 0x08080808082b082b, + 0x08080808082b2b08, + 0x08080808082b2b2b, + 0x0808080819080819, + 0x0808080819081908, + 0x0808080819190808, + 0x0808080819192b08, + 0x08080808192b0819, + 0x08080808192b1908, + 0x080808082b080808, + 0x080808082b08082b, + 0x080808082b082b2b, + 0x080808082b2b082b, + 0x0808081908080819, + 0x0808081908081908, + 0x0808081908190808, + 0x0808081908191919, + 0x0808081919080808, + 0x080808192b081908, + 0x080808192b192b08, + 0x0808082b08080808, + 0x0808082b0808082b, + 0x0808082b082b082b, + 0x0808082b2b08082b, + 0x0808190808080819, + 0x0808190808081908, + 0x0808190808190808, + 0x08081908082b0819, + 0x08081908082b1908, + 0x0808190819080808, + 0x080819081908082b, + 0x0808190819082b08, + 0x08081908192b0808, + 0x080819082b080819, + 0x080819082b081908, + 0x080819082b190808, + 0x080819082b2b1908, + 0x0808191908080808, + 0x080819190808082b, + 0x0808191908082b08, + 0x08081919082b0808, + 0x080819191908192b, + 0x08081919192b2b19, + 0x080819192b080808, + 0x080819192b190819, + 0x0808192b08082b19, + 0x0808192b08190808, + 0x0808192b19080808, + 0x0808192b2b081908, + 0x0808192b2b2b1908, + 0x08082b0808080808, + 0x08082b0808081919, + 0x08082b0808082b08, + 0x08082b0808191908, + 0x08082b08082b2b08, + 0x08082b0819080819, + 0x08082b0819081908, + 0x08082b0819190808, + 0x08082b081919082b, + 0x08082b082b082b08, + 0x08082b1908081908, + 0x08082b1919080808, + 0x08082b2b0808082b, + 0x08082b2b08191908, + 0x0819080808080819, + 0x0819080808081908, + 0x0819080808190808, + 0x08190808082b0819, + 0x0819080819080808, + 0x08190808192b0808, + 0x081908082b081908, + 0x081908082b190808, + 0x081908082b191919, + 0x0819081908080808, + 0x0819081908082b08, + 0x08190819082b0808, + 0x0819081919190808, + 0x0819081919192b2b, + 0x081908192b080808, + 0x0819082b082b1908, + 0x0819082b19081919, + 0x0819190808080808, + 0x0819190808082b08, + 0x08191908082b0808, + 0x08191908082b1919, + 0x0819190819082b19, + 0x081919082b080808, + 0x0819191908192b08, + 0x08191919192b082b, + 0x0819192b08080808, + 0x0819192b0819192b, + 0x08192b0808080819, + 0x08192b0808081908, + 0x08192b0808190808, + 0x08192b0819080808, + 0x08192b082b080819, + 0x08192b1908080808, + 0x08192b1908081919, + 0x08192b192b2b0808, + 0x08192b2b19190819, + 0x082b080808080808, + 0x082b08080808082b, + 0x082b080808082b2b, + 0x082b080819081908, + 0x082b0808192b0819, + 0x082b08082b080808, + 0x082b08082b08082b, + 0x082b0819082b2b19, + 0x082b081919082b08, + 0x082b082b08080808, + 0x082b082b0808082b, + 0x082b190808080819, + 0x082b190808081908, + 0x082b190808190808, + 0x082b190819080808, + 0x082b19081919192b, + 0x082b191908080808, + 0x082b191919080819, + 0x082b1919192b1908, + 0x082b192b2b190808, + 0x082b2b0808082b08, + 0x082b2b08082b0808, + 0x082b2b082b191908, + 0x082b2b2b19081908, + 0x1908080808080819, + 0x1908080808081908, + 0x1908080808190808, + 0x1908080808192b08, + 0x19080808082b0819, + 0x19080808082b1908, + 0x1908080819080808, + 0x1908080819082b08, + 0x190808081919192b, + 0x19080808192b0808, + 0x190808082b080819, + 0x190808082b081908, + 0x190808082b190808, + 0x1908081908080808, + 0x19080819082b0808, + 0x19080819192b0819, + 0x190808192b080808, + 0x190808192b081919, + 0x1908082b08080819, + 0x1908082b08190808, + 0x1908082b19082b08, + 0x1908082b1919192b, + 0x1908082b192b2b08, + 0x1908190808080808, + 0x1908190808082b08, + 0x19081908082b0808, + 0x190819082b080808, + 0x190819082b192b19, + 0x190819190819082b, + 0x19081919082b1908, + 0x1908192b08080808, + 0x19082b0808080819, + 0x19082b0808081908, + 0x19082b0808190808, + 0x19082b0819080808, + 0x19082b0819081919, + 0x19082b1908080808, + 0x19082b1919192b08, + 0x19082b19192b0819, + 0x19082b192b08082b, + 0x19082b2b19081919, + 0x19082b2b2b190808, + 0x1919080808080808, + 0x1919080808082b08, + 0x1919080808190819, + 0x1919080808192b19, + 0x19190808082b0808, + 0x191908082b080808, + 0x191908082b082b08, + 0x1919081908081908, + 0x191908191908082b, + 0x191908192b2b1908, + 0x1919082b2b190819, + 0x191919082b190808, + 0x191919082b19082b, + 0x1919191908082b2b, + 0x1919192b08080819, + 0x1919192b19191908, + 0x19192b0808080808, + 0x19192b0808190819, + 0x19192b0808192b19, + 0x19192b08192b1908, + 0x19192b1919080808, + 0x19192b2b08082b08, + 0x192b080808081908, + 0x192b080808190808, + 0x192b080819080808, + 0x192b0808192b2b08, + 0x192b081908080808, + 0x192b081919191919, + 0x192b082b08192b08, + 0x192b082b192b0808, + 0x192b190808080808, + 0x192b190808081919, + 0x192b191908190808, + 0x192b19190819082b, + 0x192b19192b081908, + 0x192b2b081908082b, + 0x2b08080808080808, + 0x2b0808080808082b, + 0x2b08080808082b2b, + 0x2b08080819080819, + 0x2b0808082b08082b, + 0x2b08081908081908, + 0x2b08081908192b08, + 0x2b08081919080808, + 0x2b08082b08190819, + 0x2b08190808080819, + 0x2b08190808081908, + 0x2b08190808190808, + 0x2b08190808191919, + 0x2b08190819080808, + 0x2b081908192b0808, + 0x2b08191908080808, + 0x2b0819191908192b, + 0x2b0819192b191908, + 0x2b08192b08082b19, + 0x2b08192b19080808, + 0x2b08192b192b0808, + 0x2b082b080808082b, + 0x2b082b1908081908, + 0x2b082b2b08190819, + 0x2b19080808081908, + 0x2b19080808190808, + 0x2b190808082b1908, + 0x2b19080819080808, + 0x2b1908082b2b0819, + 0x2b1908190819192b, + 0x2b1908192b080808, + 0x2b19082b19081919, + 0x2b19190808080808, + 0x2b191908082b082b, + 0x2b19190819081908, + 0x2b19191919190819, + 0x2b192b082b080819, + 0x2b192b19082b0808, + 0x2b2b08080808082b, + 0x2b2b080819190808, + 0x2b2b08082b081919, + 0x2b2b081908082b19, + 0x2b2b082b08080808, + 0x2b2b190808192b08, + 0x2b2b2b0819190808, + 0x2b2b2b1908081908, +}; +constexpr uint64_t iq2xs_grid[512] = { + 0x0808080808080808, + 0x080808080808082b, + 0x0808080808081919, + 0x0808080808082b08, + 0x0808080808082b2b, + 0x0808080808190819, + 0x0808080808191908, + 0x080808080819192b, + 0x0808080808192b19, + 0x08080808082b0808, + 0x08080808082b082b, + 0x08080808082b1919, + 0x08080808082b2b08, + 0x0808080819080819, + 0x0808080819081908, + 0x080808081908192b, + 0x0808080819082b19, + 0x0808080819190808, + 0x080808081919082b, + 0x0808080819191919, + 0x0808080819192b08, + 0x08080808192b0819, + 0x08080808192b1908, + 0x080808082b080808, + 0x080808082b08082b, + 0x080808082b081919, + 0x080808082b082b08, + 0x080808082b190819, + 0x080808082b191908, + 0x080808082b192b19, + 0x080808082b2b0808, + 0x0808081908080819, + 0x0808081908081908, + 0x080808190808192b, + 0x0808081908082b19, + 0x0808081908190808, + 0x080808190819082b, + 0x0808081908191919, + 0x0808081908192b08, + 0x0808081908192b2b, + 0x08080819082b0819, + 0x08080819082b1908, + 0x0808081919080808, + 0x080808191908082b, + 0x0808081919081919, + 0x0808081919082b08, + 0x0808081919190819, + 0x0808081919191908, + 0x08080819192b0808, + 0x08080819192b2b08, + 0x080808192b080819, + 0x080808192b081908, + 0x080808192b190808, + 0x0808082b08080808, + 0x0808082b0808082b, + 0x0808082b08081919, + 0x0808082b08082b08, + 0x0808082b08190819, + 0x0808082b08191908, + 0x0808082b082b0808, + 0x0808082b19080819, + 0x0808082b19081908, + 0x0808082b19190808, + 0x0808082b19191919, + 0x0808082b2b080808, + 0x0808082b2b082b2b, + 0x0808190808080819, + 0x0808190808081908, + 0x080819080808192b, + 0x0808190808082b19, + 0x0808190808190808, + 0x080819080819082b, + 0x0808190808191919, + 0x0808190808192b08, + 0x08081908082b0819, + 0x08081908082b1908, + 0x0808190819080808, + 0x080819081908082b, + 0x0808190819081919, + 0x0808190819082b08, + 0x0808190819190819, + 0x0808190819191908, + 0x080819081919192b, + 0x08081908192b0808, + 0x080819082b080819, + 0x080819082b081908, + 0x080819082b190808, + 0x0808191908080808, + 0x080819190808082b, + 0x0808191908081919, + 0x0808191908082b08, + 0x0808191908190819, + 0x0808191908191908, + 0x08081919082b0808, + 0x0808191919080819, + 0x0808191919081908, + 0x0808191919190808, + 0x08081919192b0819, + 0x080819192b080808, + 0x0808192b08080819, + 0x0808192b08081908, + 0x0808192b08190808, + 0x0808192b082b192b, + 0x0808192b19080808, + 0x0808192b1908082b, + 0x0808192b2b081908, + 0x08082b0808080808, + 0x08082b080808082b, + 0x08082b0808081919, + 0x08082b0808082b08, + 0x08082b0808082b2b, + 0x08082b0808190819, + 0x08082b0808191908, + 0x08082b08082b0808, + 0x08082b08082b1919, + 0x08082b0819080819, + 0x08082b0819081908, + 0x08082b0819190808, + 0x08082b0819192b08, + 0x08082b082b080808, + 0x08082b082b2b0808, + 0x08082b082b2b2b2b, + 0x08082b1908080819, + 0x08082b1908081908, + 0x08082b1908190808, + 0x08082b1919080808, + 0x08082b192b080819, + 0x08082b192b082b19, + 0x08082b2b08080808, + 0x08082b2b082b0808, + 0x08082b2b082b2b08, + 0x08082b2b2b19192b, + 0x08082b2b2b2b0808, + 0x0819080808080819, + 0x0819080808081908, + 0x081908080808192b, + 0x0819080808082b19, + 0x0819080808190808, + 0x081908080819082b, + 0x0819080808191919, + 0x0819080808192b08, + 0x08190808082b0819, + 0x08190808082b1908, + 0x0819080819080808, + 0x081908081908082b, + 0x0819080819081919, + 0x0819080819082b08, + 0x0819080819190819, + 0x0819080819191908, + 0x08190808192b0808, + 0x08190808192b2b2b, + 0x081908082b080819, + 0x081908082b081908, + 0x081908082b190808, + 0x0819081908080808, + 0x081908190808082b, + 0x0819081908081919, + 0x0819081908082b08, + 0x0819081908190819, + 0x0819081908191908, + 0x08190819082b0808, + 0x0819081919080819, + 0x0819081919081908, + 0x0819081919190808, + 0x081908192b080808, + 0x081908192b191908, + 0x081908192b19192b, + 0x0819082b08080819, + 0x0819082b08081908, + 0x0819082b0808192b, + 0x0819082b08190808, + 0x0819082b19080808, + 0x0819082b192b0808, + 0x0819190808080808, + 0x081919080808082b, + 0x0819190808081919, + 0x0819190808082b08, + 0x0819190808190819, + 0x0819190808191908, + 0x08191908082b0808, + 0x0819190819080819, + 0x0819190819081908, + 0x0819190819082b19, + 0x0819190819190808, + 0x08191908192b1908, + 0x081919082b080808, + 0x0819191908080819, + 0x0819191908081908, + 0x0819191908190808, + 0x0819191919080808, + 0x0819192b08080808, + 0x0819192b08191908, + 0x0819192b19082b19, + 0x08192b0808080819, + 0x08192b0808081908, + 0x08192b0808190808, + 0x08192b080819082b, + 0x08192b0819080808, + 0x08192b0819191908, + 0x08192b082b08192b, + 0x08192b1908080808, + 0x08192b1908081919, + 0x08192b19192b192b, + 0x08192b2b19190819, + 0x08192b2b2b2b2b19, + 0x082b080808080808, + 0x082b08080808082b, + 0x082b080808081919, + 0x082b080808082b08, + 0x082b080808082b2b, + 0x082b080808190819, + 0x082b080808191908, + 0x082b0808082b0808, + 0x082b080819080819, + 0x082b080819081908, + 0x082b080819190808, + 0x082b08082b080808, + 0x082b08082b2b0808, + 0x082b081908080819, + 0x082b081908081908, + 0x082b081908190808, + 0x082b081919080808, + 0x082b081919082b08, + 0x082b0819192b1919, + 0x082b082b08080808, + 0x082b082b082b082b, + 0x082b082b2b080808, + 0x082b082b2b2b2b08, + 0x082b190808080819, + 0x082b190808081908, + 0x082b190808190808, + 0x082b1908082b2b19, + 0x082b190819080808, + 0x082b191908080808, + 0x082b191919080819, + 0x082b19191919082b, + 0x082b19192b192b19, + 0x082b192b08080819, + 0x082b192b08192b2b, + 0x082b192b2b2b192b, + 0x082b2b0808080808, + 0x082b2b0808082b08, + 0x082b2b0808082b2b, + 0x082b2b08082b0808, + 0x082b2b0819191919, + 0x082b2b082b082b08, + 0x082b2b082b2b082b, + 0x082b2b19192b2b08, + 0x082b2b192b190808, + 0x082b2b2b08082b08, + 0x082b2b2b082b0808, + 0x082b2b2b2b08082b, + 0x082b2b2b2b082b08, + 0x082b2b2b2b082b2b, + 0x1908080808080819, + 0x1908080808081908, + 0x190808080808192b, + 0x1908080808082b19, + 0x1908080808190808, + 0x190808080819082b, + 0x1908080808191919, + 0x1908080808192b08, + 0x19080808082b0819, + 0x19080808082b1908, + 0x1908080819080808, + 0x190808081908082b, + 0x1908080819081919, + 0x1908080819082b08, + 0x1908080819082b2b, + 0x1908080819190819, + 0x1908080819191908, + 0x19080808192b0808, + 0x19080808192b1919, + 0x190808082b080819, + 0x190808082b081908, + 0x190808082b190808, + 0x1908081908080808, + 0x190808190808082b, + 0x1908081908081919, + 0x1908081908082b08, + 0x1908081908190819, + 0x1908081908191908, + 0x19080819082b0808, + 0x1908081919080819, + 0x1908081919081908, + 0x1908081919190808, + 0x190808192b080808, + 0x190808192b081919, + 0x190808192b2b082b, + 0x1908082b08080819, + 0x1908082b08081908, + 0x1908082b08190808, + 0x1908082b0819082b, + 0x1908082b082b2b19, + 0x1908082b19080808, + 0x1908190808080808, + 0x190819080808082b, + 0x1908190808081919, + 0x1908190808082b08, + 0x1908190808190819, + 0x1908190808191908, + 0x1908190808192b19, + 0x19081908082b0808, + 0x1908190819080819, + 0x1908190819081908, + 0x1908190819190808, + 0x190819082b080808, + 0x190819082b191908, + 0x1908191908080819, + 0x1908191908081908, + 0x1908191908190808, + 0x19081919082b1908, + 0x1908191919080808, + 0x190819192b192b2b, + 0x1908192b08080808, + 0x1908192b08082b2b, + 0x1908192b19081908, + 0x1908192b19190808, + 0x19082b0808080819, + 0x19082b0808081908, + 0x19082b0808190808, + 0x19082b0819080808, + 0x19082b0819081919, + 0x19082b0819191908, + 0x19082b08192b082b, + 0x19082b1908080808, + 0x19082b1908190819, + 0x19082b1919081908, + 0x19082b1919190808, + 0x19082b19192b2b19, + 0x19082b2b08081908, + 0x1919080808080808, + 0x191908080808082b, + 0x1919080808081919, + 0x1919080808082b08, + 0x1919080808190819, + 0x1919080808191908, + 0x19190808082b0808, + 0x19190808082b2b08, + 0x1919080819080819, + 0x1919080819081908, + 0x1919080819190808, + 0x191908082b080808, + 0x1919081908080819, + 0x1919081908081908, + 0x1919081908190808, + 0x1919081908191919, + 0x1919081919080808, + 0x191908191908082b, + 0x1919082b08080808, + 0x1919082b19081908, + 0x1919082b2b2b2b2b, + 0x1919190808080819, + 0x1919190808081908, + 0x1919190808190808, + 0x19191908082b0819, + 0x1919190819080808, + 0x19191908192b0808, + 0x191919082b080819, + 0x191919082b2b0819, + 0x1919191908080808, + 0x1919191908082b08, + 0x191919192b080808, + 0x191919192b082b08, + 0x1919192b082b0819, + 0x1919192b192b2b08, + 0x1919192b2b2b0819, + 0x19192b0808080808, + 0x19192b0808191908, + 0x19192b0819080819, + 0x19192b0819190808, + 0x19192b082b192b19, + 0x19192b1908192b2b, + 0x19192b1919080808, + 0x19192b191908082b, + 0x19192b2b2b081919, + 0x192b080808080819, + 0x192b080808081908, + 0x192b080808190808, + 0x192b080819080808, + 0x192b080819191908, + 0x192b0808192b082b, + 0x192b08082b08192b, + 0x192b08082b2b2b19, + 0x192b081908080808, + 0x192b082b082b1908, + 0x192b082b19082b2b, + 0x192b082b2b19082b, + 0x192b190808080808, + 0x192b19080819192b, + 0x192b191908190808, + 0x192b191919080808, + 0x192b191919081919, + 0x192b19192b2b1908, + 0x192b2b0808080819, + 0x192b2b08192b2b2b, + 0x192b2b19082b1919, + 0x192b2b2b0808192b, + 0x192b2b2b19191908, + 0x192b2b2b192b082b, + 0x2b08080808080808, + 0x2b0808080808082b, + 0x2b08080808081919, + 0x2b08080808082b08, + 0x2b08080808190819, + 0x2b08080808191908, + 0x2b080808082b0808, + 0x2b080808082b2b2b, + 0x2b08080819080819, + 0x2b08080819081908, + 0x2b08080819190808, + 0x2b0808082b080808, + 0x2b0808082b08082b, + 0x2b0808082b2b2b08, + 0x2b0808082b2b2b2b, + 0x2b08081908080819, + 0x2b08081908081908, + 0x2b0808190808192b, + 0x2b08081908190808, + 0x2b08081919080808, + 0x2b08081919190819, + 0x2b08081919192b19, + 0x2b08082b08080808, + 0x2b08082b082b0808, + 0x2b08082b2b080808, + 0x2b08082b2b08082b, + 0x2b08082b2b2b0808, + 0x2b08082b2b2b2b08, + 0x2b08190808080819, + 0x2b08190808081908, + 0x2b08190808190808, + 0x2b0819080819082b, + 0x2b08190808191919, + 0x2b08190819080808, + 0x2b081908192b0808, + 0x2b0819082b082b19, + 0x2b08191908080808, + 0x2b08191919081908, + 0x2b0819192b2b1919, + 0x2b08192b08192b08, + 0x2b08192b192b2b2b, + 0x2b082b0808080808, + 0x2b082b0808082b08, + 0x2b082b08082b1919, + 0x2b082b0819192b2b, + 0x2b082b082b080808, + 0x2b082b082b08082b, + 0x2b082b082b2b2b08, + 0x2b082b190808192b, + 0x2b082b2b082b082b, + 0x2b082b2b2b080808, + 0x2b082b2b2b082b08, + 0x2b082b2b2b19192b, + 0x2b082b2b2b2b2b08, + 0x2b19080808080819, + 0x2b19080808081908, + 0x2b19080808190808, + 0x2b19080819080808, + 0x2b1908081919192b, + 0x2b1908082b081908, + 0x2b19081908080808, + 0x2b190819082b082b, + 0x2b190819192b1908, + 0x2b19082b1919192b, + 0x2b19082b2b082b19, + 0x2b19190808080808, + 0x2b19190808081919, + 0x2b19190819081908, + 0x2b19190819190808, + 0x2b19190819192b08, + 0x2b191919082b2b19, + 0x2b1919192b190808, + 0x2b1919192b19082b, + 0x2b19192b19080819, + 0x2b192b0819190819, + 0x2b192b082b2b192b, + 0x2b192b1919082b19, + 0x2b192b2b08191919, + 0x2b192b2b192b0808, + 0x2b2b080808080808, + 0x2b2b08080808082b, + 0x2b2b080808082b08, + 0x2b2b080808082b2b, + 0x2b2b0808082b0808, + 0x2b2b0808082b2b2b, + 0x2b2b08082b2b0808, + 0x2b2b081919190819, + 0x2b2b081919192b19, + 0x2b2b08192b2b192b, + 0x2b2b082b08080808, + 0x2b2b082b0808082b, + 0x2b2b082b08082b08, + 0x2b2b082b082b2b2b, + 0x2b2b082b2b080808, + 0x2b2b082b2b2b0808, + 0x2b2b190819080808, + 0x2b2b19082b191919, + 0x2b2b192b192b1919, + 0x2b2b192b2b192b08, + 0x2b2b2b0808082b2b, + 0x2b2b2b08082b0808, + 0x2b2b2b08082b082b, + 0x2b2b2b08082b2b08, + 0x2b2b2b082b2b0808, + 0x2b2b2b082b2b2b08, + 0x2b2b2b1908081908, + 0x2b2b2b192b081908, + 0x2b2b2b192b08192b, + 0x2b2b2b2b082b2b08, + 0x2b2b2b2b082b2b2b, + 0x2b2b2b2b2b190819, + 0x2b2b2b2b2b2b2b2b, +}; +constexpr uint64_t iq2s_grid[1024] = { + 0x0808080808080808, + 0x080808080808082b, + 0x0808080808081919, + 0x0808080808082b08, + 0x0808080808082b2b, + 0x0808080808190819, + 0x0808080808191908, + 0x080808080819192b, + 0x0808080808192b19, + 0x08080808082b0808, + 0x08080808082b082b, + 0x08080808082b1919, + 0x08080808082b2b08, + 0x0808080819080819, + 0x0808080819081908, + 0x080808081908192b, + 0x0808080819082b19, + 0x0808080819190808, + 0x080808081919082b, + 0x0808080819191919, + 0x0808080819192b08, + 0x08080808192b0819, + 0x08080808192b1908, + 0x08080808192b192b, + 0x08080808192b2b19, + 0x080808082b080808, + 0x080808082b08082b, + 0x080808082b081919, + 0x080808082b082b08, + 0x080808082b190819, + 0x080808082b191908, + 0x080808082b2b0808, + 0x080808082b2b1919, + 0x080808082b2b2b2b, + 0x0808081908080819, + 0x0808081908081908, + 0x080808190808192b, + 0x0808081908082b19, + 0x0808081908190808, + 0x080808190819082b, + 0x0808081908191919, + 0x0808081908192b08, + 0x08080819082b0819, + 0x08080819082b1908, + 0x0808081919080808, + 0x080808191908082b, + 0x0808081919081919, + 0x0808081919082b08, + 0x0808081919190819, + 0x0808081919191908, + 0x080808191919192b, + 0x0808081919192b19, + 0x08080819192b0808, + 0x08080819192b1919, + 0x08080819192b2b08, + 0x080808192b080819, + 0x080808192b081908, + 0x080808192b190808, + 0x080808192b19082b, + 0x080808192b191919, + 0x080808192b2b0819, + 0x080808192b2b1908, + 0x0808082b08080808, + 0x0808082b0808082b, + 0x0808082b08081919, + 0x0808082b08082b08, + 0x0808082b08190819, + 0x0808082b08191908, + 0x0808082b082b0808, + 0x0808082b082b2b2b, + 0x0808082b19080819, + 0x0808082b19081908, + 0x0808082b1908192b, + 0x0808082b19082b19, + 0x0808082b19190808, + 0x0808082b19191919, + 0x0808082b2b080808, + 0x0808082b2b081919, + 0x0808082b2b082b2b, + 0x0808082b2b191908, + 0x0808082b2b2b082b, + 0x0808190808080819, + 0x0808190808081908, + 0x080819080808192b, + 0x0808190808082b19, + 0x0808190808190808, + 0x080819080819082b, + 0x0808190808191919, + 0x0808190808192b08, + 0x08081908082b0819, + 0x08081908082b1908, + 0x08081908082b192b, + 0x08081908082b2b19, + 0x0808190819080808, + 0x080819081908082b, + 0x0808190819081919, + 0x0808190819082b08, + 0x0808190819082b2b, + 0x0808190819190819, + 0x0808190819191908, + 0x080819081919192b, + 0x0808190819192b19, + 0x08081908192b0808, + 0x08081908192b082b, + 0x08081908192b1919, + 0x080819082b080819, + 0x080819082b081908, + 0x080819082b08192b, + 0x080819082b082b19, + 0x080819082b190808, + 0x080819082b191919, + 0x080819082b192b08, + 0x080819082b2b0819, + 0x080819082b2b1908, + 0x0808191908080808, + 0x080819190808082b, + 0x0808191908081919, + 0x0808191908082b08, + 0x0808191908082b2b, + 0x0808191908190819, + 0x0808191908191908, + 0x080819190819192b, + 0x0808191908192b19, + 0x08081919082b0808, + 0x08081919082b1919, + 0x08081919082b2b08, + 0x0808191919080819, + 0x0808191919081908, + 0x080819191908192b, + 0x0808191919082b19, + 0x0808191919190808, + 0x080819191919082b, + 0x0808191919191919, + 0x0808191919192b08, + 0x08081919192b0819, + 0x08081919192b1908, + 0x080819192b080808, + 0x080819192b08082b, + 0x080819192b081919, + 0x080819192b082b08, + 0x080819192b190819, + 0x080819192b191908, + 0x080819192b2b0808, + 0x0808192b08080819, + 0x0808192b08081908, + 0x0808192b0808192b, + 0x0808192b08082b19, + 0x0808192b08190808, + 0x0808192b08191919, + 0x0808192b19080808, + 0x0808192b19081919, + 0x0808192b19082b08, + 0x0808192b19190819, + 0x0808192b19191908, + 0x0808192b192b0808, + 0x0808192b2b080819, + 0x0808192b2b081908, + 0x0808192b2b190808, + 0x08082b0808080808, + 0x08082b080808082b, + 0x08082b0808081919, + 0x08082b0808082b08, + 0x08082b0808190819, + 0x08082b0808191908, + 0x08082b080819192b, + 0x08082b0808192b19, + 0x08082b08082b0808, + 0x08082b08082b1919, + 0x08082b08082b2b2b, + 0x08082b0819080819, + 0x08082b0819081908, + 0x08082b081908192b, + 0x08082b0819082b19, + 0x08082b0819190808, + 0x08082b081919082b, + 0x08082b0819191919, + 0x08082b0819192b08, + 0x08082b08192b0819, + 0x08082b08192b1908, + 0x08082b082b080808, + 0x08082b082b081919, + 0x08082b082b191908, + 0x08082b082b2b2b2b, + 0x08082b1908080819, + 0x08082b1908081908, + 0x08082b1908190808, + 0x08082b190819082b, + 0x08082b1908191919, + 0x08082b1908192b08, + 0x08082b19082b0819, + 0x08082b1919080808, + 0x08082b1919081919, + 0x08082b1919082b08, + 0x08082b1919190819, + 0x08082b1919191908, + 0x08082b19192b0808, + 0x08082b192b080819, + 0x08082b192b190808, + 0x08082b2b08080808, + 0x08082b2b08190819, + 0x08082b2b08191908, + 0x08082b2b082b082b, + 0x08082b2b082b2b08, + 0x08082b2b082b2b2b, + 0x08082b2b19190808, + 0x08082b2b2b192b19, + 0x0819080808080819, + 0x0819080808081908, + 0x081908080808192b, + 0x0819080808082b19, + 0x0819080808190808, + 0x081908080819082b, + 0x0819080808191919, + 0x0819080808192b08, + 0x08190808082b0819, + 0x08190808082b1908, + 0x08190808082b192b, + 0x0819080819080808, + 0x081908081908082b, + 0x0819080819081919, + 0x0819080819082b08, + 0x0819080819190819, + 0x0819080819191908, + 0x081908081919192b, + 0x0819080819192b19, + 0x08190808192b0808, + 0x08190808192b082b, + 0x08190808192b1919, + 0x08190808192b2b08, + 0x081908082b080819, + 0x081908082b081908, + 0x081908082b08192b, + 0x081908082b190808, + 0x081908082b191919, + 0x081908082b192b08, + 0x081908082b2b0819, + 0x081908082b2b1908, + 0x0819081908080808, + 0x081908190808082b, + 0x0819081908081919, + 0x0819081908082b08, + 0x0819081908082b2b, + 0x0819081908190819, + 0x0819081908191908, + 0x081908190819192b, + 0x0819081908192b19, + 0x08190819082b0808, + 0x08190819082b082b, + 0x08190819082b1919, + 0x08190819082b2b08, + 0x0819081919080819, + 0x0819081919081908, + 0x081908191908192b, + 0x0819081919082b19, + 0x0819081919190808, + 0x081908191919082b, + 0x0819081919191919, + 0x0819081919192b08, + 0x08190819192b0819, + 0x08190819192b1908, + 0x081908192b080808, + 0x081908192b08082b, + 0x081908192b081919, + 0x081908192b082b08, + 0x081908192b190819, + 0x081908192b191908, + 0x0819082b08080819, + 0x0819082b08081908, + 0x0819082b08082b19, + 0x0819082b08190808, + 0x0819082b08191919, + 0x0819082b082b0819, + 0x0819082b082b1908, + 0x0819082b19080808, + 0x0819082b19081919, + 0x0819082b19190819, + 0x0819082b19191908, + 0x0819082b2b080819, + 0x0819082b2b081908, + 0x0819082b2b190808, + 0x0819190808080808, + 0x081919080808082b, + 0x0819190808081919, + 0x0819190808082b08, + 0x0819190808190819, + 0x0819190808191908, + 0x081919080819192b, + 0x0819190808192b19, + 0x08191908082b0808, + 0x08191908082b1919, + 0x08191908082b2b08, + 0x0819190819080819, + 0x0819190819081908, + 0x081919081908192b, + 0x0819190819082b19, + 0x0819190819190808, + 0x081919081919082b, + 0x0819190819191919, + 0x0819190819192b08, + 0x08191908192b0819, + 0x08191908192b1908, + 0x081919082b080808, + 0x081919082b08082b, + 0x081919082b081919, + 0x081919082b082b08, + 0x081919082b190819, + 0x081919082b191908, + 0x081919082b2b0808, + 0x0819191908080819, + 0x0819191908081908, + 0x081919190808192b, + 0x0819191908082b19, + 0x0819191908190808, + 0x081919190819082b, + 0x0819191908191919, + 0x0819191908192b08, + 0x08191919082b0819, + 0x08191919082b1908, + 0x0819191919080808, + 0x081919191908082b, + 0x0819191919081919, + 0x0819191919082b08, + 0x0819191919190819, + 0x0819191919191908, + 0x08191919192b0808, + 0x081919192b080819, + 0x081919192b081908, + 0x081919192b190808, + 0x0819192b08080808, + 0x0819192b08081919, + 0x0819192b08082b08, + 0x0819192b08190819, + 0x0819192b08191908, + 0x0819192b082b0808, + 0x0819192b19080819, + 0x0819192b19081908, + 0x0819192b19190808, + 0x0819192b2b080808, + 0x0819192b2b2b2b2b, + 0x08192b0808080819, + 0x08192b0808081908, + 0x08192b080808192b, + 0x08192b0808082b19, + 0x08192b0808190808, + 0x08192b0808191919, + 0x08192b0808192b08, + 0x08192b08082b0819, + 0x08192b0819080808, + 0x08192b081908082b, + 0x08192b0819081919, + 0x08192b0819082b08, + 0x08192b0819190819, + 0x08192b0819191908, + 0x08192b08192b0808, + 0x08192b082b080819, + 0x08192b082b081908, + 0x08192b1908080808, + 0x08192b190808082b, + 0x08192b1908081919, + 0x08192b1908082b08, + 0x08192b1908190819, + 0x08192b1908191908, + 0x08192b19082b0808, + 0x08192b1919080819, + 0x08192b1919081908, + 0x08192b1919190808, + 0x08192b19192b2b19, + 0x08192b192b2b082b, + 0x08192b2b08081908, + 0x08192b2b08190808, + 0x08192b2b19080808, + 0x08192b2b1919192b, + 0x082b080808080808, + 0x082b08080808082b, + 0x082b080808081919, + 0x082b080808082b08, + 0x082b080808190819, + 0x082b080808191908, + 0x082b08080819192b, + 0x082b080808192b19, + 0x082b0808082b0808, + 0x082b0808082b1919, + 0x082b0808082b2b2b, + 0x082b080819080819, + 0x082b080819081908, + 0x082b080819190808, + 0x082b08081919082b, + 0x082b080819191919, + 0x082b0808192b1908, + 0x082b08082b080808, + 0x082b08082b082b2b, + 0x082b08082b191908, + 0x082b08082b2b2b2b, + 0x082b081908080819, + 0x082b081908081908, + 0x082b081908190808, + 0x082b08190819082b, + 0x082b081908191919, + 0x082b0819082b0819, + 0x082b081919080808, + 0x082b08191908082b, + 0x082b081919081919, + 0x082b081919190819, + 0x082b081919191908, + 0x082b0819192b0808, + 0x082b08192b080819, + 0x082b08192b081908, + 0x082b08192b190808, + 0x082b082b08080808, + 0x082b082b08082b2b, + 0x082b082b082b082b, + 0x082b082b082b2b08, + 0x082b082b082b2b2b, + 0x082b082b19081908, + 0x082b082b19190808, + 0x082b082b2b082b08, + 0x082b082b2b082b2b, + 0x082b082b2b2b2b08, + 0x082b190808080819, + 0x082b190808081908, + 0x082b19080808192b, + 0x082b190808082b19, + 0x082b190808190808, + 0x082b190808191919, + 0x082b190808192b08, + 0x082b1908082b0819, + 0x082b1908082b1908, + 0x082b190819080808, + 0x082b19081908082b, + 0x082b190819081919, + 0x082b190819082b08, + 0x082b190819190819, + 0x082b190819191908, + 0x082b1908192b0808, + 0x082b19082b080819, + 0x082b19082b081908, + 0x082b19082b190808, + 0x082b191908080808, + 0x082b191908081919, + 0x082b191908082b08, + 0x082b191908190819, + 0x082b191908191908, + 0x082b1919082b0808, + 0x082b191919080819, + 0x082b191919081908, + 0x082b191919190808, + 0x082b1919192b192b, + 0x082b19192b080808, + 0x082b192b08080819, + 0x082b192b08081908, + 0x082b192b08190808, + 0x082b192b19080808, + 0x082b192b19192b19, + 0x082b2b0808080808, + 0x082b2b0808081919, + 0x082b2b0808190819, + 0x082b2b0808191908, + 0x082b2b0819080819, + 0x082b2b0819081908, + 0x082b2b0819190808, + 0x082b2b082b082b2b, + 0x082b2b082b2b2b2b, + 0x082b2b1908080819, + 0x082b2b1908081908, + 0x082b2b1908190808, + 0x082b2b192b191919, + 0x082b2b2b08082b2b, + 0x082b2b2b082b082b, + 0x082b2b2b192b1908, + 0x082b2b2b2b082b08, + 0x082b2b2b2b082b2b, + 0x1908080808080819, + 0x1908080808081908, + 0x190808080808192b, + 0x1908080808082b19, + 0x1908080808190808, + 0x190808080819082b, + 0x1908080808191919, + 0x1908080808192b08, + 0x1908080808192b2b, + 0x19080808082b0819, + 0x19080808082b1908, + 0x19080808082b192b, + 0x1908080819080808, + 0x190808081908082b, + 0x1908080819081919, + 0x1908080819082b08, + 0x1908080819082b2b, + 0x1908080819190819, + 0x1908080819191908, + 0x190808081919192b, + 0x1908080819192b19, + 0x19080808192b0808, + 0x19080808192b082b, + 0x19080808192b1919, + 0x190808082b080819, + 0x190808082b081908, + 0x190808082b190808, + 0x190808082b191919, + 0x190808082b192b08, + 0x190808082b2b0819, + 0x190808082b2b1908, + 0x1908081908080808, + 0x190808190808082b, + 0x1908081908081919, + 0x1908081908082b08, + 0x1908081908190819, + 0x1908081908191908, + 0x190808190819192b, + 0x1908081908192b19, + 0x19080819082b0808, + 0x19080819082b082b, + 0x19080819082b1919, + 0x1908081919080819, + 0x1908081919081908, + 0x190808191908192b, + 0x1908081919082b19, + 0x1908081919190808, + 0x190808191919082b, + 0x1908081919191919, + 0x1908081919192b08, + 0x19080819192b0819, + 0x19080819192b1908, + 0x190808192b080808, + 0x190808192b08082b, + 0x190808192b081919, + 0x190808192b082b08, + 0x190808192b190819, + 0x190808192b191908, + 0x190808192b2b0808, + 0x1908082b08080819, + 0x1908082b08081908, + 0x1908082b08190808, + 0x1908082b0819082b, + 0x1908082b08191919, + 0x1908082b08192b08, + 0x1908082b082b1908, + 0x1908082b19080808, + 0x1908082b19081919, + 0x1908082b19082b08, + 0x1908082b19190819, + 0x1908082b19191908, + 0x1908082b192b0808, + 0x1908082b2b080819, + 0x1908082b2b081908, + 0x1908190808080808, + 0x190819080808082b, + 0x1908190808081919, + 0x1908190808082b08, + 0x1908190808082b2b, + 0x1908190808190819, + 0x1908190808191908, + 0x190819080819192b, + 0x1908190808192b19, + 0x19081908082b0808, + 0x19081908082b082b, + 0x19081908082b1919, + 0x19081908082b2b08, + 0x1908190819080819, + 0x1908190819081908, + 0x190819081908192b, + 0x1908190819082b19, + 0x1908190819190808, + 0x190819081919082b, + 0x1908190819191919, + 0x1908190819192b08, + 0x19081908192b0819, + 0x19081908192b1908, + 0x190819082b080808, + 0x190819082b08082b, + 0x190819082b081919, + 0x190819082b082b08, + 0x190819082b190819, + 0x190819082b191908, + 0x190819082b2b0808, + 0x1908191908080819, + 0x1908191908081908, + 0x190819190808192b, + 0x1908191908082b19, + 0x1908191908190808, + 0x190819190819082b, + 0x1908191908191919, + 0x1908191908192b08, + 0x19081919082b0819, + 0x19081919082b1908, + 0x1908191919080808, + 0x190819191908082b, + 0x1908191919081919, + 0x1908191919082b08, + 0x1908191919190819, + 0x1908191919191908, + 0x19081919192b0808, + 0x19081919192b2b2b, + 0x190819192b080819, + 0x190819192b081908, + 0x190819192b190808, + 0x1908192b08080808, + 0x1908192b0808082b, + 0x1908192b08081919, + 0x1908192b08082b08, + 0x1908192b08190819, + 0x1908192b08191908, + 0x1908192b082b0808, + 0x1908192b19080819, + 0x1908192b19081908, + 0x1908192b19190808, + 0x1908192b2b080808, + 0x1908192b2b2b1919, + 0x19082b0808080819, + 0x19082b0808081908, + 0x19082b0808082b19, + 0x19082b0808190808, + 0x19082b080819082b, + 0x19082b0808191919, + 0x19082b0808192b08, + 0x19082b08082b0819, + 0x19082b08082b1908, + 0x19082b0819080808, + 0x19082b081908082b, + 0x19082b0819081919, + 0x19082b0819082b08, + 0x19082b0819190819, + 0x19082b0819191908, + 0x19082b08192b0808, + 0x19082b082b081908, + 0x19082b082b190808, + 0x19082b1908080808, + 0x19082b190808082b, + 0x19082b1908081919, + 0x19082b1908082b08, + 0x19082b1908190819, + 0x19082b1908191908, + 0x19082b19082b0808, + 0x19082b1919080819, + 0x19082b1919081908, + 0x19082b1919190808, + 0x19082b192b080808, + 0x19082b192b19192b, + 0x19082b2b08080819, + 0x19082b2b08081908, + 0x19082b2b08190808, + 0x19082b2b19080808, + 0x1919080808080808, + 0x191908080808082b, + 0x1919080808081919, + 0x1919080808082b08, + 0x1919080808190819, + 0x1919080808191908, + 0x191908080819192b, + 0x1919080808192b19, + 0x19190808082b0808, + 0x19190808082b082b, + 0x19190808082b1919, + 0x19190808082b2b08, + 0x1919080819080819, + 0x1919080819081908, + 0x191908081908192b, + 0x1919080819082b19, + 0x1919080819190808, + 0x191908081919082b, + 0x1919080819191919, + 0x1919080819192b08, + 0x19190808192b0819, + 0x19190808192b1908, + 0x191908082b080808, + 0x191908082b08082b, + 0x191908082b081919, + 0x191908082b082b08, + 0x191908082b190819, + 0x191908082b191908, + 0x1919081908080819, + 0x1919081908081908, + 0x191908190808192b, + 0x1919081908082b19, + 0x1919081908190808, + 0x191908190819082b, + 0x1919081908191919, + 0x1919081908192b08, + 0x19190819082b0819, + 0x19190819082b1908, + 0x1919081919080808, + 0x191908191908082b, + 0x1919081919081919, + 0x1919081919082b08, + 0x1919081919190819, + 0x1919081919191908, + 0x19190819192b0808, + 0x191908192b080819, + 0x191908192b081908, + 0x191908192b190808, + 0x1919082b08080808, + 0x1919082b08081919, + 0x1919082b08082b08, + 0x1919082b08190819, + 0x1919082b08191908, + 0x1919082b082b0808, + 0x1919082b19080819, + 0x1919082b19081908, + 0x1919082b19190808, + 0x1919082b192b2b19, + 0x1919082b2b080808, + 0x1919190808080819, + 0x1919190808081908, + 0x191919080808192b, + 0x1919190808082b19, + 0x1919190808190808, + 0x191919080819082b, + 0x1919190808191919, + 0x1919190808192b08, + 0x19191908082b0819, + 0x19191908082b1908, + 0x1919190819080808, + 0x191919081908082b, + 0x1919190819081919, + 0x1919190819082b08, + 0x1919190819190819, + 0x1919190819191908, + 0x19191908192b0808, + 0x191919082b080819, + 0x191919082b081908, + 0x191919082b190808, + 0x1919191908080808, + 0x191919190808082b, + 0x1919191908081919, + 0x1919191908082b08, + 0x1919191908190819, + 0x1919191908191908, + 0x19191919082b0808, + 0x1919191919080819, + 0x1919191919081908, + 0x1919191919190808, + 0x191919192b080808, + 0x1919192b08080819, + 0x1919192b08081908, + 0x1919192b08190808, + 0x1919192b082b192b, + 0x1919192b19080808, + 0x19192b0808080808, + 0x19192b080808082b, + 0x19192b0808081919, + 0x19192b0808082b08, + 0x19192b0808190819, + 0x19192b0808191908, + 0x19192b08082b0808, + 0x19192b0819080819, + 0x19192b0819081908, + 0x19192b0819190808, + 0x19192b0819192b2b, + 0x19192b082b080808, + 0x19192b1908080819, + 0x19192b1908081908, + 0x19192b1908190808, + 0x19192b1919080808, + 0x19192b2b08080808, + 0x19192b2b08192b19, + 0x19192b2b2b081919, + 0x19192b2b2b2b2b08, + 0x192b080808080819, + 0x192b080808081908, + 0x192b08080808192b, + 0x192b080808190808, + 0x192b08080819082b, + 0x192b080808191919, + 0x192b080808192b08, + 0x192b0808082b0819, + 0x192b0808082b1908, + 0x192b080819080808, + 0x192b080819081919, + 0x192b080819082b08, + 0x192b080819190819, + 0x192b080819191908, + 0x192b0808192b0808, + 0x192b08082b081908, + 0x192b08082b190808, + 0x192b081908080808, + 0x192b08190808082b, + 0x192b081908081919, + 0x192b081908082b08, + 0x192b081908190819, + 0x192b081908191908, + 0x192b0819082b0808, + 0x192b081919080819, + 0x192b081919081908, + 0x192b081919190808, + 0x192b08192b080808, + 0x192b08192b192b19, + 0x192b082b08081908, + 0x192b082b08190808, + 0x192b082b19080808, + 0x192b082b1919192b, + 0x192b082b2b2b0819, + 0x192b190808080808, + 0x192b190808081919, + 0x192b190808082b08, + 0x192b190808190819, + 0x192b190808191908, + 0x192b1908082b0808, + 0x192b190819080819, + 0x192b190819081908, + 0x192b190819190808, + 0x192b19082b080808, + 0x192b191908080819, + 0x192b191908081908, + 0x192b191908190808, + 0x192b191919080808, + 0x192b191919082b2b, + 0x192b1919192b2b08, + 0x192b19192b19082b, + 0x192b192b08080808, + 0x192b192b2b191908, + 0x192b2b0808080819, + 0x192b2b0808081908, + 0x192b2b0808190808, + 0x192b2b08192b1919, + 0x192b2b082b192b08, + 0x192b2b1908080808, + 0x192b2b19082b2b2b, + 0x192b2b2b1908082b, + 0x192b2b2b2b2b0819, + 0x2b08080808080808, + 0x2b0808080808082b, + 0x2b08080808081919, + 0x2b08080808082b08, + 0x2b08080808190819, + 0x2b08080808191908, + 0x2b08080808192b19, + 0x2b080808082b0808, + 0x2b080808082b1919, + 0x2b08080819080819, + 0x2b08080819081908, + 0x2b08080819190808, + 0x2b0808081919082b, + 0x2b08080819191919, + 0x2b08080819192b08, + 0x2b080808192b0819, + 0x2b0808082b080808, + 0x2b0808082b081919, + 0x2b0808082b190819, + 0x2b0808082b191908, + 0x2b08081908080819, + 0x2b08081908081908, + 0x2b08081908082b19, + 0x2b08081908190808, + 0x2b0808190819082b, + 0x2b08081908191919, + 0x2b08081908192b08, + 0x2b080819082b0819, + 0x2b080819082b1908, + 0x2b08081919080808, + 0x2b0808191908082b, + 0x2b08081919081919, + 0x2b08081919082b08, + 0x2b08081919190819, + 0x2b08081919191908, + 0x2b0808192b080819, + 0x2b0808192b081908, + 0x2b0808192b190808, + 0x2b0808192b2b2b19, + 0x2b08082b08080808, + 0x2b08082b08081919, + 0x2b08082b08082b2b, + 0x2b08082b08190819, + 0x2b08082b08191908, + 0x2b08082b19080819, + 0x2b08082b19081908, + 0x2b08082b19190808, + 0x2b08190808080819, + 0x2b08190808081908, + 0x2b0819080808192b, + 0x2b08190808082b19, + 0x2b08190808190808, + 0x2b0819080819082b, + 0x2b08190808191919, + 0x2b08190808192b08, + 0x2b081908082b0819, + 0x2b08190819080808, + 0x2b0819081908082b, + 0x2b08190819081919, + 0x2b08190819082b08, + 0x2b08190819190819, + 0x2b08190819191908, + 0x2b081908192b0808, + 0x2b0819082b080819, + 0x2b0819082b081908, + 0x2b0819082b190808, + 0x2b08191908080808, + 0x2b0819190808082b, + 0x2b08191908081919, + 0x2b08191908082b08, + 0x2b08191908190819, + 0x2b08191908191908, + 0x2b081919082b0808, + 0x2b08191919080819, + 0x2b08191919081908, + 0x2b08191919190808, + 0x2b0819192b080808, + 0x2b0819192b082b2b, + 0x2b08192b08080819, + 0x2b08192b08081908, + 0x2b08192b08190808, + 0x2b08192b082b2b19, + 0x2b08192b19080808, + 0x2b082b0808080808, + 0x2b082b0808081919, + 0x2b082b0808190819, + 0x2b082b0808191908, + 0x2b082b0819080819, + 0x2b082b0819081908, + 0x2b082b0819190808, + 0x2b082b082b2b082b, + 0x2b082b1908080819, + 0x2b082b1908081908, + 0x2b082b1919080808, + 0x2b082b19192b1919, + 0x2b082b2b082b082b, + 0x2b082b2b19192b08, + 0x2b082b2b19192b2b, + 0x2b082b2b2b08082b, + 0x2b082b2b2b2b082b, + 0x2b19080808080819, + 0x2b19080808081908, + 0x2b19080808082b19, + 0x2b19080808190808, + 0x2b1908080819082b, + 0x2b19080808191919, + 0x2b19080808192b08, + 0x2b190808082b1908, + 0x2b19080819080808, + 0x2b1908081908082b, + 0x2b19080819081919, + 0x2b19080819082b08, + 0x2b19080819190819, + 0x2b19080819191908, + 0x2b190808192b0808, + 0x2b1908082b080819, + 0x2b1908082b081908, + 0x2b1908082b190808, + 0x2b19081908080808, + 0x2b19081908081919, + 0x2b19081908190819, + 0x2b19081908191908, + 0x2b19081919080819, + 0x2b19081919081908, + 0x2b19081919190808, + 0x2b19081919192b2b, + 0x2b19082b08080819, + 0x2b19082b08081908, + 0x2b19082b08190808, + 0x2b19082b19080808, + 0x2b19082b2b2b192b, + 0x2b19190808080808, + 0x2b1919080808082b, + 0x2b19190808081919, + 0x2b19190808082b08, + 0x2b19190808190819, + 0x2b19190808191908, + 0x2b191908082b0808, + 0x2b19190819080819, + 0x2b19190819081908, + 0x2b19190819190808, + 0x2b1919082b080808, + 0x2b1919082b19192b, + 0x2b19191908080819, + 0x2b19191908081908, + 0x2b19191908190808, + 0x2b19191919080808, + 0x2b1919192b192b08, + 0x2b1919192b2b0819, + 0x2b19192b08080808, + 0x2b19192b1908192b, + 0x2b19192b192b1908, + 0x2b192b0808080819, + 0x2b192b0808081908, + 0x2b192b0808190808, + 0x2b192b08082b192b, + 0x2b192b0819080808, + 0x2b192b082b2b2b19, + 0x2b192b1908080808, + 0x2b192b1919082b19, + 0x2b192b191919082b, + 0x2b192b2b2b190808, + 0x2b2b080808080808, + 0x2b2b080808081919, + 0x2b2b080808082b2b, + 0x2b2b080808191908, + 0x2b2b0808082b082b, + 0x2b2b0808082b2b2b, + 0x2b2b080819080819, + 0x2b2b080819081908, + 0x2b2b080819190808, + 0x2b2b08082b2b082b, + 0x2b2b08082b2b2b2b, + 0x2b2b081919080808, + 0x2b2b0819192b1919, + 0x2b2b082b0808082b, + 0x2b2b082b08082b2b, + 0x2b2b082b082b082b, + 0x2b2b082b082b2b08, + 0x2b2b082b082b2b2b, + 0x2b2b082b2b08082b, + 0x2b2b082b2b082b08, + 0x2b2b082b2b082b2b, + 0x2b2b082b2b2b2b08, + 0x2b2b190808080819, + 0x2b2b190808081908, + 0x2b2b190808190808, + 0x2b2b190819080808, + 0x2b2b19082b082b19, + 0x2b2b19082b2b1908, + 0x2b2b191908080808, + 0x2b2b191908192b19, + 0x2b2b192b19190819, + 0x2b2b2b0808082b2b, + 0x2b2b2b08082b2b08, + 0x2b2b2b082b2b082b, + 0x2b2b2b1919191908, + 0x2b2b2b192b08192b, + 0x2b2b2b2b08082b08, + 0x2b2b2b2b08082b2b, + 0x2b2b2b2b082b0808, + 0x2b2b2b2b082b082b, + 0x2b2b2b2b082b2b08, + 0x2b2b2b2b2b082b08, + 0x2b2b2b2b2b2b2b2b, +}; +constexpr uint32_t iq3xxs_grid[256] = { + 0x04040404, + 0x04040414, + 0x04040424, + 0x04040c0c, + 0x04040c1c, + 0x04040c3e, + 0x04041404, + 0x04041414, + 0x04041c0c, + 0x04042414, + 0x04043e1c, + 0x04043e2c, + 0x040c040c, + 0x040c041c, + 0x040c0c04, + 0x040c0c14, + 0x040c140c, + 0x040c142c, + 0x040c1c04, + 0x040c1c14, + 0x040c240c, + 0x040c2c24, + 0x040c3e04, + 0x04140404, + 0x04140414, + 0x04140424, + 0x04140c0c, + 0x04141404, + 0x04141414, + 0x04141c0c, + 0x04141c1c, + 0x04141c3e, + 0x04142c0c, + 0x04142c3e, + 0x04143e2c, + 0x041c040c, + 0x041c043e, + 0x041c0c04, + 0x041c0c14, + 0x041c142c, + 0x041c3e04, + 0x04240c1c, + 0x04241c3e, + 0x04242424, + 0x04242c3e, + 0x04243e1c, + 0x04243e2c, + 0x042c040c, + 0x042c043e, + 0x042c1c14, + 0x042c2c14, + 0x04341c2c, + 0x04343424, + 0x043e0c04, + 0x043e0c24, + 0x043e0c34, + 0x043e241c, + 0x043e340c, + 0x0c04040c, + 0x0c04041c, + 0x0c040c04, + 0x0c040c14, + 0x0c04140c, + 0x0c04141c, + 0x0c041c04, + 0x0c041c14, + 0x0c041c24, + 0x0c04243e, + 0x0c042c04, + 0x0c0c0404, + 0x0c0c0414, + 0x0c0c0c0c, + 0x0c0c1404, + 0x0c0c1414, + 0x0c14040c, + 0x0c14041c, + 0x0c140c04, + 0x0c140c14, + 0x0c14140c, + 0x0c141c04, + 0x0c143e14, + 0x0c1c0404, + 0x0c1c0414, + 0x0c1c1404, + 0x0c1c1c0c, + 0x0c1c2434, + 0x0c1c3434, + 0x0c24040c, + 0x0c24042c, + 0x0c242c04, + 0x0c2c1404, + 0x0c2c1424, + 0x0c2c2434, + 0x0c2c3e0c, + 0x0c34042c, + 0x0c3e1414, + 0x0c3e2404, + 0x14040404, + 0x14040414, + 0x14040c0c, + 0x14040c1c, + 0x14041404, + 0x14041414, + 0x14041434, + 0x14041c0c, + 0x14042414, + 0x140c040c, + 0x140c041c, + 0x140c042c, + 0x140c0c04, + 0x140c0c14, + 0x140c140c, + 0x140c1c04, + 0x140c341c, + 0x140c343e, + 0x140c3e04, + 0x14140404, + 0x14140414, + 0x14140c0c, + 0x14140c3e, + 0x14141404, + 0x14141414, + 0x14141c3e, + 0x14142404, + 0x14142c2c, + 0x141c040c, + 0x141c0c04, + 0x141c0c24, + 0x141c3e04, + 0x141c3e24, + 0x14241c2c, + 0x14242c1c, + 0x142c041c, + 0x142c143e, + 0x142c240c, + 0x142c3e24, + 0x143e040c, + 0x143e041c, + 0x143e0c34, + 0x143e242c, + 0x1c04040c, + 0x1c040c04, + 0x1c040c14, + 0x1c04140c, + 0x1c04141c, + 0x1c042c04, + 0x1c04342c, + 0x1c043e14, + 0x1c0c0404, + 0x1c0c0414, + 0x1c0c1404, + 0x1c0c1c0c, + 0x1c0c2424, + 0x1c0c2434, + 0x1c14040c, + 0x1c14041c, + 0x1c140c04, + 0x1c14142c, + 0x1c142c14, + 0x1c143e14, + 0x1c1c0c0c, + 0x1c1c1c1c, + 0x1c241c04, + 0x1c24243e, + 0x1c243e14, + 0x1c2c0404, + 0x1c2c0434, + 0x1c2c1414, + 0x1c2c2c2c, + 0x1c340c24, + 0x1c341c34, + 0x1c34341c, + 0x1c3e1c1c, + 0x1c3e3404, + 0x24040424, + 0x24040c3e, + 0x24041c2c, + 0x24041c3e, + 0x24042c1c, + 0x24042c3e, + 0x240c3e24, + 0x24141404, + 0x24141c3e, + 0x24142404, + 0x24143404, + 0x24143434, + 0x241c043e, + 0x241c242c, + 0x24240424, + 0x24242c0c, + 0x24243424, + 0x242c142c, + 0x242c241c, + 0x242c3e04, + 0x243e042c, + 0x243e0c04, + 0x243e0c14, + 0x243e1c04, + 0x2c040c14, + 0x2c04240c, + 0x2c043e04, + 0x2c0c0404, + 0x2c0c0434, + 0x2c0c1434, + 0x2c0c2c2c, + 0x2c140c24, + 0x2c141c14, + 0x2c143e14, + 0x2c1c0414, + 0x2c1c2c1c, + 0x2c240c04, + 0x2c24141c, + 0x2c24143e, + 0x2c243e14, + 0x2c2c0414, + 0x2c2c1c0c, + 0x2c342c04, + 0x2c3e1424, + 0x2c3e2414, + 0x34041424, + 0x34042424, + 0x34042434, + 0x34043424, + 0x340c140c, + 0x340c340c, + 0x34140c3e, + 0x34143424, + 0x341c1c04, + 0x341c1c34, + 0x34242424, + 0x342c042c, + 0x342c2c14, + 0x34341c1c, + 0x343e041c, + 0x343e140c, + 0x3e04041c, + 0x3e04042c, + 0x3e04043e, + 0x3e040c04, + 0x3e041c14, + 0x3e042c14, + 0x3e0c1434, + 0x3e0c2404, + 0x3e140c14, + 0x3e14242c, + 0x3e142c14, + 0x3e1c0404, + 0x3e1c0c2c, + 0x3e1c1c1c, + 0x3e1c3404, + 0x3e24140c, + 0x3e24240c, + 0x3e2c0404, + 0x3e2c0414, + 0x3e2c1424, + 0x3e341c04, +}; +constexpr uint32_t iq3s_grid[512] = { + 0x01010101, + 0x01010103, + 0x01010105, + 0x0101010b, + 0x0101010f, + 0x01010301, + 0x01010303, + 0x01010305, + 0x01010309, + 0x0101030d, + 0x01010501, + 0x01010503, + 0x0101050b, + 0x01010707, + 0x01010901, + 0x01010905, + 0x0101090b, + 0x0101090f, + 0x01010b03, + 0x01010b07, + 0x01010d01, + 0x01010d05, + 0x01010f03, + 0x01010f09, + 0x01010f0f, + 0x01030101, + 0x01030103, + 0x01030105, + 0x01030109, + 0x01030301, + 0x01030303, + 0x0103030b, + 0x01030501, + 0x01030507, + 0x0103050f, + 0x01030703, + 0x0103070b, + 0x01030909, + 0x01030d03, + 0x01030d0b, + 0x01030f05, + 0x01050101, + 0x01050103, + 0x0105010b, + 0x0105010f, + 0x01050301, + 0x01050307, + 0x0105030d, + 0x01050503, + 0x0105050b, + 0x01050701, + 0x01050709, + 0x01050905, + 0x0105090b, + 0x0105090f, + 0x01050b03, + 0x01050b07, + 0x01050f01, + 0x01050f07, + 0x01070107, + 0x01070303, + 0x0107030b, + 0x01070501, + 0x01070505, + 0x01070703, + 0x01070707, + 0x0107070d, + 0x01070909, + 0x01070b01, + 0x01070b05, + 0x01070d0f, + 0x01070f03, + 0x01070f0b, + 0x01090101, + 0x01090307, + 0x0109030f, + 0x01090503, + 0x01090509, + 0x01090705, + 0x01090901, + 0x01090907, + 0x01090b03, + 0x01090f01, + 0x010b0105, + 0x010b0109, + 0x010b0501, + 0x010b0505, + 0x010b050d, + 0x010b0707, + 0x010b0903, + 0x010b090b, + 0x010b090f, + 0x010b0d0d, + 0x010b0f07, + 0x010d010d, + 0x010d0303, + 0x010d0307, + 0x010d0703, + 0x010d0b05, + 0x010d0f03, + 0x010f0101, + 0x010f0105, + 0x010f0109, + 0x010f0501, + 0x010f0505, + 0x010f050d, + 0x010f0707, + 0x010f0b01, + 0x010f0b09, + 0x03010101, + 0x03010103, + 0x03010105, + 0x03010109, + 0x03010301, + 0x03010303, + 0x03010307, + 0x0301030b, + 0x0301030f, + 0x03010501, + 0x03010505, + 0x03010703, + 0x03010709, + 0x0301070d, + 0x03010b09, + 0x03010b0d, + 0x03010d03, + 0x03010f05, + 0x03030101, + 0x03030103, + 0x03030107, + 0x0303010d, + 0x03030301, + 0x03030309, + 0x03030503, + 0x03030701, + 0x03030707, + 0x03030903, + 0x03030b01, + 0x03030b05, + 0x03030f01, + 0x03030f0d, + 0x03050101, + 0x03050305, + 0x0305030b, + 0x0305030f, + 0x03050501, + 0x03050509, + 0x03050705, + 0x03050901, + 0x03050907, + 0x03050b0b, + 0x03050d01, + 0x03050f05, + 0x03070103, + 0x03070109, + 0x0307010f, + 0x03070301, + 0x03070307, + 0x03070503, + 0x0307050f, + 0x03070701, + 0x03070709, + 0x03070903, + 0x03070d05, + 0x03070f01, + 0x03090107, + 0x0309010b, + 0x03090305, + 0x03090309, + 0x03090703, + 0x03090707, + 0x03090905, + 0x0309090d, + 0x03090b01, + 0x03090b09, + 0x030b0103, + 0x030b0301, + 0x030b0307, + 0x030b0503, + 0x030b0701, + 0x030b0705, + 0x030b0b03, + 0x030d0501, + 0x030d0509, + 0x030d050f, + 0x030d0909, + 0x030d090d, + 0x030f0103, + 0x030f0107, + 0x030f0301, + 0x030f0305, + 0x030f0503, + 0x030f070b, + 0x030f0903, + 0x030f0d05, + 0x030f0f01, + 0x05010101, + 0x05010103, + 0x05010107, + 0x0501010b, + 0x0501010f, + 0x05010301, + 0x05010305, + 0x05010309, + 0x0501030d, + 0x05010503, + 0x05010507, + 0x0501050f, + 0x05010701, + 0x05010705, + 0x05010903, + 0x05010907, + 0x0501090b, + 0x05010b01, + 0x05010b05, + 0x05010d0f, + 0x05010f01, + 0x05010f07, + 0x05010f0b, + 0x05030101, + 0x05030105, + 0x05030301, + 0x05030307, + 0x0503030f, + 0x05030505, + 0x0503050b, + 0x05030703, + 0x05030709, + 0x05030905, + 0x05030b03, + 0x05050103, + 0x05050109, + 0x0505010f, + 0x05050503, + 0x05050507, + 0x05050701, + 0x0505070f, + 0x05050903, + 0x05050b07, + 0x05050b0f, + 0x05050f03, + 0x05050f09, + 0x05070101, + 0x05070105, + 0x0507010b, + 0x05070303, + 0x05070505, + 0x05070509, + 0x05070703, + 0x05070707, + 0x05070905, + 0x05070b01, + 0x05070d0d, + 0x05090103, + 0x0509010f, + 0x05090501, + 0x05090507, + 0x05090705, + 0x0509070b, + 0x05090903, + 0x05090f05, + 0x05090f0b, + 0x050b0109, + 0x050b0303, + 0x050b0505, + 0x050b070f, + 0x050b0901, + 0x050b0b07, + 0x050b0f01, + 0x050d0101, + 0x050d0105, + 0x050d010f, + 0x050d0503, + 0x050d0b0b, + 0x050d0d03, + 0x050f010b, + 0x050f0303, + 0x050f050d, + 0x050f0701, + 0x050f0907, + 0x050f0b01, + 0x07010105, + 0x07010303, + 0x07010307, + 0x0701030b, + 0x0701030f, + 0x07010505, + 0x07010703, + 0x07010707, + 0x0701070b, + 0x07010905, + 0x07010909, + 0x0701090f, + 0x07010b03, + 0x07010d07, + 0x07010f03, + 0x07030103, + 0x07030107, + 0x0703010b, + 0x07030309, + 0x07030503, + 0x07030507, + 0x07030901, + 0x07030d01, + 0x07030f05, + 0x07030f0d, + 0x07050101, + 0x07050305, + 0x07050501, + 0x07050705, + 0x07050709, + 0x07050b01, + 0x07070103, + 0x07070301, + 0x07070309, + 0x07070503, + 0x07070507, + 0x0707050f, + 0x07070701, + 0x07070903, + 0x07070907, + 0x0707090f, + 0x07070b0b, + 0x07070f07, + 0x07090107, + 0x07090303, + 0x0709030d, + 0x07090505, + 0x07090703, + 0x07090b05, + 0x07090d01, + 0x07090d09, + 0x070b0103, + 0x070b0301, + 0x070b0305, + 0x070b050b, + 0x070b0705, + 0x070b0909, + 0x070b0b0d, + 0x070b0f07, + 0x070d030d, + 0x070d0903, + 0x070f0103, + 0x070f0107, + 0x070f0501, + 0x070f0505, + 0x070f070b, + 0x09010101, + 0x09010109, + 0x09010305, + 0x09010501, + 0x09010509, + 0x0901050f, + 0x09010705, + 0x09010903, + 0x09010b01, + 0x09010f01, + 0x09030105, + 0x0903010f, + 0x09030303, + 0x09030307, + 0x09030505, + 0x09030701, + 0x0903070b, + 0x09030907, + 0x09030b03, + 0x09030b0b, + 0x09050103, + 0x09050107, + 0x09050301, + 0x0905030b, + 0x09050503, + 0x09050707, + 0x09050901, + 0x09050b0f, + 0x09050d05, + 0x09050f01, + 0x09070109, + 0x09070303, + 0x09070307, + 0x09070501, + 0x09070505, + 0x09070703, + 0x0907070b, + 0x09090101, + 0x09090105, + 0x09090509, + 0x0909070f, + 0x09090901, + 0x09090f03, + 0x090b010b, + 0x090b010f, + 0x090b0503, + 0x090b0d05, + 0x090d0307, + 0x090d0709, + 0x090d0d01, + 0x090f0301, + 0x090f030b, + 0x090f0701, + 0x090f0907, + 0x090f0b03, + 0x0b010105, + 0x0b010301, + 0x0b010309, + 0x0b010505, + 0x0b010901, + 0x0b010909, + 0x0b01090f, + 0x0b010b05, + 0x0b010d0d, + 0x0b010f09, + 0x0b030103, + 0x0b030107, + 0x0b03010b, + 0x0b030305, + 0x0b030503, + 0x0b030705, + 0x0b030f05, + 0x0b050101, + 0x0b050303, + 0x0b050507, + 0x0b050701, + 0x0b05070d, + 0x0b050b07, + 0x0b070105, + 0x0b07010f, + 0x0b070301, + 0x0b07050f, + 0x0b070909, + 0x0b070b03, + 0x0b070d0b, + 0x0b070f07, + 0x0b090103, + 0x0b090109, + 0x0b090501, + 0x0b090705, + 0x0b09090d, + 0x0b0b0305, + 0x0b0b050d, + 0x0b0b0b03, + 0x0b0b0b07, + 0x0b0d0905, + 0x0b0f0105, + 0x0b0f0109, + 0x0b0f0505, + 0x0d010303, + 0x0d010307, + 0x0d01030b, + 0x0d010703, + 0x0d010707, + 0x0d010d01, + 0x0d030101, + 0x0d030501, + 0x0d03050f, + 0x0d030d09, + 0x0d050305, + 0x0d050709, + 0x0d050905, + 0x0d050b0b, + 0x0d050d05, + 0x0d050f01, + 0x0d070101, + 0x0d070309, + 0x0d070503, + 0x0d070901, + 0x0d09050b, + 0x0d090907, + 0x0d090d05, + 0x0d0b0101, + 0x0d0b0107, + 0x0d0b0709, + 0x0d0b0d01, + 0x0d0d010b, + 0x0d0d0901, + 0x0d0f0303, + 0x0d0f0307, + 0x0f010101, + 0x0f010109, + 0x0f01010f, + 0x0f010501, + 0x0f010505, + 0x0f01070d, + 0x0f010901, + 0x0f010b09, + 0x0f010d05, + 0x0f030105, + 0x0f030303, + 0x0f030509, + 0x0f030907, + 0x0f03090b, + 0x0f050103, + 0x0f050109, + 0x0f050301, + 0x0f05030d, + 0x0f050503, + 0x0f050701, + 0x0f050b03, + 0x0f070105, + 0x0f070705, + 0x0f07070b, + 0x0f070b07, + 0x0f090103, + 0x0f09010b, + 0x0f090307, + 0x0f090501, + 0x0f090b01, + 0x0f0b0505, + 0x0f0b0905, + 0x0f0d0105, + 0x0f0d0703, + 0x0f0f0101, +}; +constexpr uint64_t iq1s_grid[2048] = { + 0xffffffffffffffff, + 0xffffffffffffff01, + 0xffffffffffff0000, + 0xffffffffffff01ff, + 0xffffffffffff0101, + 0xffffffffff00ff00, + 0xffffffffff000000, + 0xffffffffff01ffff, + 0xffffffffff01ff01, + 0xffffffffff0101ff, + 0xffffffffff010101, + 0xffffffff00ff0000, + 0xffffffff0000ff00, + 0xffffffff000000ff, + 0xffffffff00000001, + 0xffffffff00010000, + 0xffffffff01ffffff, + 0xffffffff01ffff01, + 0xffffffff01ff01ff, + 0xffffffff01ff0101, + 0xffffffff01000000, + 0xffffffff0101ffff, + 0xffffffff0101ff01, + 0xffffffff010101ff, + 0xffffffff01010101, + 0xffffff00ffff00ff, + 0xffffff00ffff0000, + 0xffffff00ff00ff00, + 0xffffff00ff0000ff, + 0xffffff00ff000001, + 0xffffff00ff000100, + 0xffffff00ff000101, + 0xffffff00ff010000, + 0xffffff0000ffff00, + 0xffffff0000ff0001, + 0xffffff0000ff0100, + 0xffffff000000ff01, + 0xffffff0000000000, + 0xffffff0000000101, + 0xffffff000001ff00, + 0xffffff00000100ff, + 0xffffff0000010001, + 0xffffff00000101ff, + 0xffffff0001ff0000, + 0xffffff000100ff00, + 0xffffff00010000ff, + 0xffffff0001000001, + 0xffffff0001010000, + 0xffffff01ffffffff, + 0xffffff01ffffff01, + 0xffffff01ffff01ff, + 0xffffff01ffff0101, + 0xffffff01ff000000, + 0xffffff01ff01ffff, + 0xffffff01ff01ff01, + 0xffffff01ff0101ff, + 0xffffff01ff010101, + 0xffffff0100ff0000, + 0xffffff010000ff00, + 0xffffff0100000100, + 0xffffff01000100ff, + 0xffffff0100010100, + 0xffffff0101ffffff, + 0xffffff0101ffff01, + 0xffffff0101ff01ff, + 0xffffff0101ff0101, + 0xffffff010100ff00, + 0xffffff0101000000, + 0xffffff0101000100, + 0xffffff010101ffff, + 0xffffff010101ff01, + 0xffffff01010101ff, + 0xffffff0101010101, + 0xffff00ffff00ff00, + 0xffff00ffff0000ff, + 0xffff00ffff000001, + 0xffff00ffff010000, + 0xffff00ff00ffff00, + 0xffff00ff00ff0100, + 0xffff00ff00000000, + 0xffff00ff00000101, + 0xffff00ff000100ff, + 0xffff00ff00010000, + 0xffff00ff0100ff00, + 0xffff00ff01000100, + 0xffff00ff01010000, + 0xffff0000ffffff00, + 0xffff0000ffff00ff, + 0xffff0000ffff0000, + 0xffff0000ffff0001, + 0xffff0000ff000000, + 0xffff0000ff0001ff, + 0xffff0000ff000101, + 0xffff0000ff010100, + 0xffff000000ffffff, + 0xffff000000ff0000, + 0xffff000000ff0101, + 0xffff00000000ffff, + 0xffff00000000ff00, + 0xffff0000000000ff, + 0xffff000000000000, + 0xffff000000000001, + 0xffff000000000100, + 0xffff00000001ffff, + 0xffff00000001ff01, + 0xffff000000010000, + 0xffff0000000101ff, + 0xffff000000010101, + 0xffff000001ffff00, + 0xffff00000100ff00, + 0xffff000001000000, + 0xffff0000010001ff, + 0xffff000001000101, + 0xffff00000101ff00, + 0xffff0000010100ff, + 0xffff000001010000, + 0xffff000001010001, + 0xffff000001010100, + 0xffff0001ff0000ff, + 0xffff0001ff000100, + 0xffff000100ffff00, + 0xffff000100ff00ff, + 0xffff00010000ffff, + 0xffff00010000ff01, + 0xffff000100000000, + 0xffff0001000001ff, + 0xffff00010001ffff, + 0xffff00010001ff00, + 0xffff000100010001, + 0xffff000100010100, + 0xffff000101ff0000, + 0xffff00010100ff00, + 0xffff0001010000ff, + 0xffff000101000100, + 0xffff01ffffffffff, + 0xffff01ffffffff01, + 0xffff01ffffff01ff, + 0xffff01ffffff0101, + 0xffff01ffff000000, + 0xffff01ffff01ffff, + 0xffff01ffff01ff01, + 0xffff01ffff0101ff, + 0xffff01ffff010101, + 0xffff01ff00ff0000, + 0xffff01ff0000ff00, + 0xffff01ff00000001, + 0xffff01ff00010000, + 0xffff01ff01ffffff, + 0xffff01ff01ffff01, + 0xffff01ff01ff01ff, + 0xffff01ff01ff0101, + 0xffff01ff01000000, + 0xffff01ff0101ffff, + 0xffff01ff0101ff01, + 0xffff01ff010101ff, + 0xffff01ff01010101, + 0xffff0100ffff0000, + 0xffff0100ff00ff00, + 0xffff0100ff0000ff, + 0xffff0100ff000100, + 0xffff0100ff0100ff, + 0xffff0100ff010000, + 0xffff010000ffff00, + 0xffff01000000ffff, + 0xffff01000000ff00, + 0xffff010000000000, + 0xffff01000001ff00, + 0xffff0100000100ff, + 0xffff010000010100, + 0xffff01000100ff00, + 0xffff0100010000ff, + 0xffff010001000001, + 0xffff010001000100, + 0xffff010001010000, + 0xffff0101ffffffff, + 0xffff0101ffffff01, + 0xffff0101ffff01ff, + 0xffff0101ffff0101, + 0xffff0101ff000000, + 0xffff0101ff01ffff, + 0xffff0101ff01ff01, + 0xffff0101ff0101ff, + 0xffff0101ff010101, + 0xffff010100ff0000, + 0xffff01010000ff00, + 0xffff010100000100, + 0xffff01010001ff00, + 0xffff010100010000, + 0xffff010101ffffff, + 0xffff010101ffff01, + 0xffff010101ff0000, + 0xffff010101ff01ff, + 0xffff010101ff0101, + 0xffff010101000000, + 0xffff01010101ffff, + 0xffff01010101ff01, + 0xffff0101010101ff, + 0xffff010101010101, + 0xff00ffffff00ffff, + 0xff00ffffff00ff00, + 0xff00ffffff0000ff, + 0xff00ffffff000100, + 0xff00ffffff0100ff, + 0xff00ffffff010000, + 0xff00ffff00ffff00, + 0xff00ffff00ff00ff, + 0xff00ffff0000ffff, + 0xff00ffff00000000, + 0xff00ffff000001ff, + 0xff00ffff0001ff00, + 0xff00ffff000100ff, + 0xff00ffff00010000, + 0xff00ffff00010100, + 0xff00ffff0100ff00, + 0xff00ffff010000ff, + 0xff00ffff01000001, + 0xff00ffff0101ff00, + 0xff00ffff01010000, + 0xff00ff00ffffff00, + 0xff00ff00ffff00ff, + 0xff00ff00ffff0001, + 0xff00ff00ffff0100, + 0xff00ff00ff00ffff, + 0xff00ff00ff00ff01, + 0xff00ff00ff000000, + 0xff00ff00ff0001ff, + 0xff00ff00ff01ff00, + 0xff00ff00ff0100ff, + 0xff00ff00ff010100, + 0xff00ff0000ff0000, + 0xff00ff0000ff0101, + 0xff00ff000000ffff, + 0xff00ff000000ff00, + 0xff00ff000000ff01, + 0xff00ff00000000ff, + 0xff00ff0000000000, + 0xff00ff0000000001, + 0xff00ff0000000100, + 0xff00ff000001ffff, + 0xff00ff0000010000, + 0xff00ff0001ff00ff, + 0xff00ff000100ff01, + 0xff00ff0001000000, + 0xff00ff000101ff00, + 0xff00ff00010100ff, + 0xff00ff01ff00ff00, + 0xff00ff01ff0000ff, + 0xff00ff01ff000001, + 0xff00ff01ff010000, + 0xff00ff0100ffffff, + 0xff00ff0100ff0001, + 0xff00ff0100ff0100, + 0xff00ff010000ff01, + 0xff00ff0100000000, + 0xff00ff01000001ff, + 0xff00ff0100000101, + 0xff00ff01000100ff, + 0xff00ff0100010001, + 0xff00ff0101ff0000, + 0xff00ff010100ff00, + 0xff00ff01010000ff, + 0xff00ff0101000001, + 0xff00ff0101010000, + 0xff0000ffffffff00, + 0xff0000ffffff0001, + 0xff0000ffffff0100, + 0xff0000ffff0000ff, + 0xff0000ffff000000, + 0xff0000ffff0001ff, + 0xff0000ffff000100, + 0xff0000ffff01ff00, + 0xff0000ffff010001, + 0xff0000ff00ffff00, + 0xff0000ff00ff0000, + 0xff0000ff00ff0001, + 0xff0000ff00ff01ff, + 0xff0000ff00ff0101, + 0xff0000ff0000ff00, + 0xff0000ff000000ff, + 0xff0000ff00000000, + 0xff0000ff00000001, + 0xff0000ff00000100, + 0xff0000ff0001ff01, + 0xff0000ff00010000, + 0xff0000ff000101ff, + 0xff0000ff01ff00ff, + 0xff0000ff01ff0100, + 0xff0000ff0100ffff, + 0xff0000ff010000ff, + 0xff0000ff01000000, + 0xff0000ff010001ff, + 0xff0000ff01000100, + 0xff0000ff01000101, + 0xff0000ff0101ff00, + 0xff0000ff010100ff, + 0xff0000ff01010000, + 0xff0000ff01010100, + 0xff000000ffffff01, + 0xff000000ffff0000, + 0xff000000ffff0101, + 0xff000000ff00ff00, + 0xff000000ff0000ff, + 0xff000000ff000000, + 0xff000000ff000001, + 0xff000000ff000100, + 0xff000000ff01ffff, + 0xff000000ff01ff01, + 0xff000000ff010000, + 0xff000000ff0101ff, + 0xff000000ff010101, + 0xff00000000ffff00, + 0xff00000000ff00ff, + 0xff00000000ff0000, + 0xff00000000ff0001, + 0xff0000000000ff00, + 0xff0000000000ff01, + 0xff000000000000ff, + 0xff00000000000000, + 0xff00000000000001, + 0xff00000000000100, + 0xff00000000000101, + 0xff0000000001ff00, + 0xff000000000100ff, + 0xff00000000010000, + 0xff00000000010001, + 0xff00000000010100, + 0xff00000001ffffff, + 0xff00000001ffff01, + 0xff00000001ff00ff, + 0xff00000001ff0000, + 0xff00000001ff01ff, + 0xff00000001ff0101, + 0xff0000000100ffff, + 0xff0000000100ff00, + 0xff000000010000ff, + 0xff00000001000000, + 0xff00000001000001, + 0xff00000001000100, + 0xff00000001000101, + 0xff0000000101ffff, + 0xff0000000101ff01, + 0xff00000001010000, + 0xff000001ffffff00, + 0xff000001ffff00ff, + 0xff000001ffff0000, + 0xff000001ffff0001, + 0xff000001ff000000, + 0xff000001ff000001, + 0xff000001ff0001ff, + 0xff000001ff000101, + 0xff000001ff01ff00, + 0xff000001ff010001, + 0xff00000100ffffff, + 0xff00000100ffff01, + 0xff00000100ff00ff, + 0xff00000100ff0000, + 0xff00000100ff01ff, + 0xff00000100ff0101, + 0xff0000010000ff00, + 0xff00000100000000, + 0xff00000100000001, + 0xff000001000001ff, + 0xff00000100000100, + 0xff0000010001ff00, + 0xff000001000100ff, + 0xff00000100010000, + 0xff000001000101ff, + 0xff00000100010100, + 0xff00000100010101, + 0xff00000101ff0001, + 0xff00000101ff0101, + 0xff0000010100ff01, + 0xff00000101000000, + 0xff000001010100ff, + 0xff00000101010100, + 0xff0001ffff00ff00, + 0xff0001ffff000001, + 0xff0001ffff010000, + 0xff0001ff00ffff00, + 0xff0001ff00ff00ff, + 0xff0001ff00ff0001, + 0xff0001ff00ff0100, + 0xff0001ff0000ffff, + 0xff0001ff00000000, + 0xff0001ff000001ff, + 0xff0001ff00000101, + 0xff0001ff0001ffff, + 0xff0001ff0001ff00, + 0xff0001ff000100ff, + 0xff0001ff00010001, + 0xff0001ff00010100, + 0xff0001ff01ff0000, + 0xff0001ff0100ff00, + 0xff0001ff010000ff, + 0xff0001ff01010000, + 0xff000100ff00ffff, + 0xff000100ff00ff01, + 0xff000100ff000000, + 0xff000100ff000101, + 0xff000100ff01ff00, + 0xff000100ff010000, + 0xff00010000ffff01, + 0xff00010000ff00ff, + 0xff00010000ff0000, + 0xff00010000ff01ff, + 0xff0001000000ff00, + 0xff000100000000ff, + 0xff00010000000000, + 0xff00010000000001, + 0xff00010000000100, + 0xff00010000000101, + 0xff0001000001ffff, + 0xff00010000010000, + 0xff00010000010101, + 0xff00010001ff0100, + 0xff0001000100ff00, + 0xff0001000100ff01, + 0xff00010001000000, + 0xff000100010001ff, + 0xff0001000101ff00, + 0xff00010001010001, + 0xff00010001010100, + 0xff000101ffff0100, + 0xff000101ff000001, + 0xff000101ff0100ff, + 0xff000101ff010001, + 0xff00010100ff00ff, + 0xff00010100ff0001, + 0xff00010100ff0100, + 0xff0001010000ffff, + 0xff0001010000ff01, + 0xff00010100000000, + 0xff000101000001ff, + 0xff0001010001ff00, + 0xff00010100010001, + 0xff00010100010100, + 0xff00010101ff0000, + 0xff0001010100ff00, + 0xff00010101000001, + 0xff00010101000101, + 0xff01ffffffffffff, + 0xff01ffffffffff01, + 0xff01ffffffff01ff, + 0xff01ffffffff0101, + 0xff01ffffff000000, + 0xff01ffffff01ffff, + 0xff01ffffff01ff01, + 0xff01ffffff010000, + 0xff01ffffff0101ff, + 0xff01ffffff010101, + 0xff01ffff00ff0000, + 0xff01ffff0000ff00, + 0xff01ffff00000100, + 0xff01ffff0001ff00, + 0xff01ffff00010000, + 0xff01ffff01ffffff, + 0xff01ffff01ffff01, + 0xff01ffff01ff01ff, + 0xff01ffff01ff0101, + 0xff01ffff01000000, + 0xff01ffff0101ffff, + 0xff01ffff0101ff01, + 0xff01ffff01010000, + 0xff01ffff010101ff, + 0xff01ffff01010101, + 0xff01ff00ffff0000, + 0xff01ff00ff00ff00, + 0xff01ff00ff0000ff, + 0xff01ff00ff000100, + 0xff01ff00ff010000, + 0xff01ff0000ffff01, + 0xff01ff0000ff00ff, + 0xff01ff0000ff0100, + 0xff01ff0000000000, + 0xff01ff00000001ff, + 0xff01ff0000000101, + 0xff01ff000001ff00, + 0xff01ff00000100ff, + 0xff01ff0000010000, + 0xff01ff0000010001, + 0xff01ff0001ff0000, + 0xff01ff000100ffff, + 0xff01ff0001000001, + 0xff01ff0001000100, + 0xff01ff0001010000, + 0xff01ff01ffffff00, + 0xff01ff01ffff01ff, + 0xff01ff01ffff0101, + 0xff01ff01ff00ff00, + 0xff01ff01ff000000, + 0xff01ff01ff01ffff, + 0xff01ff01ff01ff01, + 0xff01ff01ff0101ff, + 0xff01ff01ff010101, + 0xff01ff0100ff0000, + 0xff01ff010000ff00, + 0xff01ff0100000001, + 0xff01ff0100000100, + 0xff01ff0100010000, + 0xff01ff0101ffff00, + 0xff01ff0101ff01ff, + 0xff01ff0101ff0101, + 0xff01ff010100ff00, + 0xff01ff0101000000, + 0xff01ff010101ffff, + 0xff01ff010101ff01, + 0xff01ff01010101ff, + 0xff01ff0101010101, + 0xff0100ffffff0000, + 0xff0100ffff0000ff, + 0xff0100ffff000001, + 0xff0100ffff000100, + 0xff0100ffff010000, + 0xff0100ff00ff00ff, + 0xff0100ff00ff0000, + 0xff0100ff00ff0001, + 0xff0100ff00ff0100, + 0xff0100ff0000ff01, + 0xff0100ff00000000, + 0xff0100ff000001ff, + 0xff0100ff00000101, + 0xff0100ff00010001, + 0xff0100ff01ff0000, + 0xff0100ff0100ff00, + 0xff0100ff010000ff, + 0xff0100ff01000100, + 0xff0100ff0101ff00, + 0xff0100ff01010000, + 0xff010000ffff0100, + 0xff010000ff000000, + 0xff010000ff01ff00, + 0xff010000ff010100, + 0xff01000000ffffff, + 0xff01000000ff0000, + 0xff01000000ff01ff, + 0xff0100000000ff00, + 0xff010000000000ff, + 0xff01000000000000, + 0xff01000000000100, + 0xff0100000001ff01, + 0xff01000000010000, + 0xff010000000101ff, + 0xff01000001ff0100, + 0xff0100000100ffff, + 0xff010000010000ff, + 0xff01000001000000, + 0xff010000010001ff, + 0xff01000001000101, + 0xff0100000101ff00, + 0xff010000010100ff, + 0xff01000001010001, + 0xff01000001010100, + 0xff010001ffff0000, + 0xff010001ff00ffff, + 0xff010001ff00ff01, + 0xff010001ff000100, + 0xff010001ff010000, + 0xff01000100ffff00, + 0xff01000100ff0100, + 0xff01000100000000, + 0xff0100010001ffff, + 0xff0100010001ff00, + 0xff01000100010100, + 0xff01000101ff00ff, + 0xff01000101ff0001, + 0xff0100010100ffff, + 0xff01000101000101, + 0xff0101ffffffffff, + 0xff0101ffffffff01, + 0xff0101ffffff01ff, + 0xff0101ffffff0101, + 0xff0101ffff000000, + 0xff0101ffff01ffff, + 0xff0101ffff01ff01, + 0xff0101ffff0101ff, + 0xff0101ffff010101, + 0xff0101ff00ff0000, + 0xff0101ff0000ff00, + 0xff0101ff000000ff, + 0xff0101ff00010000, + 0xff0101ff01ffffff, + 0xff0101ff01ffff01, + 0xff0101ff01ff01ff, + 0xff0101ff01ff0101, + 0xff0101ff0101ffff, + 0xff0101ff0101ff01, + 0xff0101ff010101ff, + 0xff0101ff01010101, + 0xff010100ffff0100, + 0xff010100ff00ff00, + 0xff010100ff0000ff, + 0xff010100ff000100, + 0xff010100ff010000, + 0xff01010000ff0001, + 0xff01010000ff0100, + 0xff0101000000ff01, + 0xff01010000000000, + 0xff0101000001ff00, + 0xff010100000100ff, + 0xff01010000010001, + 0xff01010000010100, + 0xff01010001ff0000, + 0xff0101000100ffff, + 0xff01010001000001, + 0xff01010001000100, + 0xff010100010100ff, + 0xff01010001010000, + 0xff010101ffffffff, + 0xff010101ffffff01, + 0xff010101ffff01ff, + 0xff010101ffff0101, + 0xff010101ff01ffff, + 0xff010101ff01ff01, + 0xff010101ff0101ff, + 0xff010101ff010101, + 0xff01010100ff0000, + 0xff0101010000ff00, + 0xff01010100000001, + 0xff01010100000100, + 0xff01010100010000, + 0xff01010101ffffff, + 0xff01010101ffff01, + 0xff01010101ff01ff, + 0xff01010101ff0101, + 0xff01010101000000, + 0xff0101010101ffff, + 0xff0101010101ff01, + 0xff010101010101ff, + 0xff01010101010101, + 0x00ffffffffff0000, + 0x00ffffffff00ff00, + 0x00ffffffff000001, + 0x00ffffffff010000, + 0x00ffffff00ff0100, + 0x00ffffff0000ff01, + 0x00ffffff00000000, + 0x00ffffff000001ff, + 0x00ffffff00000101, + 0x00ffffff0001ff00, + 0x00ffffff000100ff, + 0x00ffffff00010001, + 0x00ffffff010000ff, + 0x00ffffff01000100, + 0x00ffffff0101ff00, + 0x00ffffff01010001, + 0x00ffff00ffffffff, + 0x00ffff00ffffff00, + 0x00ffff00ffff00ff, + 0x00ffff00ffff0001, + 0x00ffff00ffff0100, + 0x00ffff00ff00ff01, + 0x00ffff00ff000000, + 0x00ffff00ff000001, + 0x00ffff00ff0001ff, + 0x00ffff00ff000101, + 0x00ffff00ff01ff00, + 0x00ffff00ff010001, + 0x00ffff00ff010100, + 0x00ffff0000ff0000, + 0x00ffff0000ff01ff, + 0x00ffff0000ff0101, + 0x00ffff000000ff00, + 0x00ffff00000000ff, + 0x00ffff0000000000, + 0x00ffff0000000001, + 0x00ffff0000000100, + 0x00ffff0000000101, + 0x00ffff0000010000, + 0x00ffff00000101ff, + 0x00ffff0000010101, + 0x00ffff0001ffff00, + 0x00ffff0001ff00ff, + 0x00ffff0001ff0001, + 0x00ffff000100ffff, + 0x00ffff000100ff01, + 0x00ffff0001000000, + 0x00ffff000101ffff, + 0x00ffff000101ff00, + 0x00ffff000101ff01, + 0x00ffff01ffff0000, + 0x00ffff01ff00ff00, + 0x00ffff01ff0000ff, + 0x00ffff01ff000001, + 0x00ffff01ff010000, + 0x00ffff0100ffff00, + 0x00ffff010000ff01, + 0x00ffff0100000000, + 0x00ffff0100000101, + 0x00ffff01000100ff, + 0x00ffff0100010100, + 0x00ffff0101ff0100, + 0x00ffff01010000ff, + 0x00ffff0101010000, + 0x00ff00ffffffff00, + 0x00ff00ffff000000, + 0x00ff00ffff000100, + 0x00ff00ffff010100, + 0x00ff00ff00ff0000, + 0x00ff00ff00ff01ff, + 0x00ff00ff00ff0101, + 0x00ff00ff0000ff00, + 0x00ff00ff000000ff, + 0x00ff00ff00000000, + 0x00ff00ff00000001, + 0x00ff00ff0001ff00, + 0x00ff00ff0001ff01, + 0x00ff00ff00010000, + 0x00ff00ff000101ff, + 0x00ff00ff00010101, + 0x00ff00ff01ffff00, + 0x00ff00ff01ff0001, + 0x00ff00ff01ff0100, + 0x00ff00ff0100ffff, + 0x00ff00ff0100ff01, + 0x00ff00ff01000000, + 0x00ff00ff0101ffff, + 0x00ff00ff0101ff00, + 0x00ff00ff01010100, + 0x00ff0000ffffff00, + 0x00ff0000ffffff01, + 0x00ff0000ffff0000, + 0x00ff0000ffff0101, + 0x00ff0000ff00ff00, + 0x00ff0000ff0000ff, + 0x00ff0000ff000000, + 0x00ff0000ff000001, + 0x00ff0000ff000100, + 0x00ff0000ff01ffff, + 0x00ff0000ff010000, + 0x00ff0000ff010101, + 0x00ff000000ffff00, + 0x00ff000000ff00ff, + 0x00ff000000ff0000, + 0x00ff000000ff0001, + 0x00ff000000ff0100, + 0x00ff00000000ffff, + 0x00ff00000000ff00, + 0x00ff0000000000ff, + 0x00ff000000000000, + 0x00ff000000000001, + 0x00ff0000000001ff, + 0x00ff000000000100, + 0x00ff00000001ff00, + 0x00ff0000000100ff, + 0x00ff000000010000, + 0x00ff000000010001, + 0x00ff000000010100, + 0x00ff000001ffff01, + 0x00ff000001ff00ff, + 0x00ff000001ff0000, + 0x00ff000001ff01ff, + 0x00ff00000100ff00, + 0x00ff0000010000ff, + 0x00ff000001000000, + 0x00ff000001000001, + 0x00ff000001000100, + 0x00ff000001000101, + 0x00ff000001010000, + 0x00ff0000010101ff, + 0x00ff000001010101, + 0x00ff0001ffffff00, + 0x00ff0001ffff0000, + 0x00ff0001ffff0100, + 0x00ff0001ff0000ff, + 0x00ff0001ff000000, + 0x00ff0001ff0001ff, + 0x00ff0001ff000101, + 0x00ff0001ff01ff00, + 0x00ff0001ff0100ff, + 0x00ff0001ff010100, + 0x00ff000100ffffff, + 0x00ff000100ffff01, + 0x00ff000100ff0000, + 0x00ff000100ff01ff, + 0x00ff00010000ffff, + 0x00ff00010000ff00, + 0x00ff00010000ff01, + 0x00ff000100000000, + 0x00ff000100000001, + 0x00ff000100000100, + 0x00ff00010001ff01, + 0x00ff000100010000, + 0x00ff0001000101ff, + 0x00ff000101ffff00, + 0x00ff000101ff0000, + 0x00ff000101ff0101, + 0x00ff0001010000ff, + 0x00ff000101000000, + 0x00ff00010101ff00, + 0x00ff0001010100ff, + 0x00ff000101010001, + 0x00ff01ffffff0000, + 0x00ff01ffff00ff00, + 0x00ff01ffff000000, + 0x00ff01ffff000101, + 0x00ff01ffff010000, + 0x00ff01ff00ffff01, + 0x00ff01ff00ff0100, + 0x00ff01ff0000ffff, + 0x00ff01ff00000000, + 0x00ff01ff000001ff, + 0x00ff01ff0001ff00, + 0x00ff01ff000100ff, + 0x00ff01ff00010001, + 0x00ff01ff00010100, + 0x00ff01ff01ff0000, + 0x00ff01ff0100ff00, + 0x00ff01ff010000ff, + 0x00ff01ff01000001, + 0x00ff01ff01000100, + 0x00ff01ff01010000, + 0x00ff0100ffffff00, + 0x00ff0100ffff0000, + 0x00ff0100ffff0001, + 0x00ff0100ffff0101, + 0x00ff0100ff00ffff, + 0x00ff0100ff0000ff, + 0x00ff0100ff000000, + 0x00ff0100ff0001ff, + 0x00ff0100ff01ff00, + 0x00ff0100ff0100ff, + 0x00ff0100ff010001, + 0x00ff010000ffffff, + 0x00ff010000ff0000, + 0x00ff010000ff0101, + 0x00ff01000000ff00, + 0x00ff01000000ff01, + 0x00ff0100000000ff, + 0x00ff010000000000, + 0x00ff010000000001, + 0x00ff010000000100, + 0x00ff01000001ffff, + 0x00ff01000001ff01, + 0x00ff010000010000, + 0x00ff010000010001, + 0x00ff010000010101, + 0x00ff010001ff0001, + 0x00ff010001ff0100, + 0x00ff01000100ff01, + 0x00ff010001000000, + 0x00ff010001000001, + 0x00ff0100010001ff, + 0x00ff01000101ff00, + 0x00ff0100010100ff, + 0x00ff010001010001, + 0x00ff010001010100, + 0x00ff0101ff000001, + 0x00ff010100ff00ff, + 0x00ff010100ff0001, + 0x00ff010100ff0100, + 0x00ff010100000000, + 0x00ff0101000001ff, + 0x00ff010100000101, + 0x00ff0101000100ff, + 0x00ff010100010100, + 0x00ff0101010000ff, + 0x00ff010101010000, + 0x0000ffffffffff00, + 0x0000ffffffff00ff, + 0x0000ffffffff0000, + 0x0000ffffffff0001, + 0x0000ffffffff0100, + 0x0000ffffff00ff01, + 0x0000ffffff000000, + 0x0000ffffff000101, + 0x0000ffffff01ff00, + 0x0000ffffff0100ff, + 0x0000ffffff010100, + 0x0000ffff00ffffff, + 0x0000ffff00ff0000, + 0x0000ffff00ff01ff, + 0x0000ffff0000ff00, + 0x0000ffff000000ff, + 0x0000ffff00000000, + 0x0000ffff00000001, + 0x0000ffff00000100, + 0x0000ffff00010000, + 0x0000ffff000101ff, + 0x0000ffff01ff0001, + 0x0000ffff01ff0100, + 0x0000ffff01000000, + 0x0000ffff010001ff, + 0x0000ffff0101ffff, + 0x0000ffff0101ff00, + 0x0000ffff01010001, + 0x0000ffff01010100, + 0x0000ff00ffff0000, + 0x0000ff00ffff01ff, + 0x0000ff00ffff0100, + 0x0000ff00ffff0101, + 0x0000ff00ff00ff00, + 0x0000ff00ff0000ff, + 0x0000ff00ff000000, + 0x0000ff00ff000001, + 0x0000ff00ff0001ff, + 0x0000ff00ff000100, + 0x0000ff00ff01ffff, + 0x0000ff00ff010000, + 0x0000ff00ff010001, + 0x0000ff00ff0101ff, + 0x0000ff00ff010101, + 0x0000ff0000ffff00, + 0x0000ff0000ff00ff, + 0x0000ff0000ff0000, + 0x0000ff0000ff0001, + 0x0000ff0000ff0100, + 0x0000ff000000ffff, + 0x0000ff000000ff00, + 0x0000ff000000ff01, + 0x0000ff00000000ff, + 0x0000ff0000000000, + 0x0000ff0000000001, + 0x0000ff00000001ff, + 0x0000ff0000000100, + 0x0000ff0000000101, + 0x0000ff000001ff00, + 0x0000ff00000100ff, + 0x0000ff0000010000, + 0x0000ff0000010001, + 0x0000ff0000010100, + 0x0000ff0001ffff01, + 0x0000ff0001ff0000, + 0x0000ff000100ff00, + 0x0000ff00010000ff, + 0x0000ff0001000000, + 0x0000ff0001000001, + 0x0000ff0001000100, + 0x0000ff000101ffff, + 0x0000ff0001010000, + 0x0000ff0001010101, + 0x0000ff01ffffff00, + 0x0000ff01ffff0001, + 0x0000ff01ff00ff01, + 0x0000ff01ff000000, + 0x0000ff01ff000101, + 0x0000ff01ff01ff00, + 0x0000ff01ff0100ff, + 0x0000ff0100ffff01, + 0x0000ff0100ff0000, + 0x0000ff0100ff0101, + 0x0000ff010000ff00, + 0x0000ff01000000ff, + 0x0000ff0100000000, + 0x0000ff0100000001, + 0x0000ff0100000100, + 0x0000ff010001ff01, + 0x0000ff0100010000, + 0x0000ff0101ff0000, + 0x0000ff010100ffff, + 0x0000ff010100ff01, + 0x0000ff0101000000, + 0x0000ff0101000100, + 0x0000ff0101000101, + 0x0000ff01010100ff, + 0x000000ffffff00ff, + 0x000000ffffff0000, + 0x000000ffff00ff00, + 0x000000ffff0000ff, + 0x000000ffff000000, + 0x000000ffff000001, + 0x000000ffff0001ff, + 0x000000ffff000100, + 0x000000ffff01ff00, + 0x000000ffff010000, + 0x000000ffff0101ff, + 0x000000ffff010101, + 0x000000ff00ffff00, + 0x000000ff00ff00ff, + 0x000000ff00ff0000, + 0x000000ff00ff0001, + 0x000000ff00ff0100, + 0x000000ff00ff0101, + 0x000000ff0000ffff, + 0x000000ff0000ff00, + 0x000000ff000000ff, + 0x000000ff00000000, + 0x000000ff00000001, + 0x000000ff000001ff, + 0x000000ff00000100, + 0x000000ff00000101, + 0x000000ff0001ff00, + 0x000000ff0001ff01, + 0x000000ff000100ff, + 0x000000ff00010000, + 0x000000ff00010001, + 0x000000ff00010100, + 0x000000ff01ffffff, + 0x000000ff01ff01ff, + 0x000000ff01ff0101, + 0x000000ff0100ff00, + 0x000000ff010000ff, + 0x000000ff01000000, + 0x000000ff01000001, + 0x000000ff01000100, + 0x000000ff0101ff00, + 0x000000ff010100ff, + 0x000000ff01010000, + 0x000000ff01010101, + 0x00000000ffffff00, + 0x00000000ffffff01, + 0x00000000ffff00ff, + 0x00000000ffff0000, + 0x00000000ffff0001, + 0x00000000ffff0100, + 0x00000000ff00ffff, + 0x00000000ff00ff00, + 0x00000000ff00ff01, + 0x00000000ff0000ff, + 0x00000000ff000000, + 0x00000000ff000001, + 0x00000000ff000100, + 0x00000000ff000101, + 0x00000000ff01ff00, + 0x00000000ff0100ff, + 0x00000000ff010000, + 0x00000000ff010001, + 0x00000000ff010100, + 0x0000000000ffffff, + 0x0000000000ffff00, + 0x0000000000ffff01, + 0x0000000000ff00ff, + 0x0000000000ff0000, + 0x0000000000ff0001, + 0x0000000000ff01ff, + 0x0000000000ff0100, + 0x000000000000ffff, + 0x000000000000ff00, + 0x000000000000ff01, + 0x00000000000000ff, + 0x0000000000000000, + 0x0000000000000001, + 0x00000000000001ff, + 0x0000000000000100, + 0x0000000000000101, + 0x000000000001ffff, + 0x000000000001ff00, + 0x00000000000100ff, + 0x0000000000010000, + 0x0000000000010001, + 0x00000000000101ff, + 0x0000000000010100, + 0x0000000000010101, + 0x0000000001ffff00, + 0x0000000001ff00ff, + 0x0000000001ff0000, + 0x0000000001ff0100, + 0x0000000001ff0101, + 0x000000000100ffff, + 0x000000000100ff00, + 0x00000000010000ff, + 0x0000000001000000, + 0x0000000001000001, + 0x00000000010001ff, + 0x0000000001000100, + 0x000000000101ff00, + 0x00000000010100ff, + 0x0000000001010000, + 0x0000000001010001, + 0x0000000001010100, + 0x00000001ffffffff, + 0x00000001ffffff00, + 0x00000001ffffff01, + 0x00000001ffff00ff, + 0x00000001ffff0001, + 0x00000001ffff01ff, + 0x00000001ffff0100, + 0x00000001ff00ff00, + 0x00000001ff0000ff, + 0x00000001ff000000, + 0x00000001ff0001ff, + 0x00000001ff000100, + 0x00000001ff01ffff, + 0x00000001ff01ff00, + 0x00000001ff01ff01, + 0x00000001ff0100ff, + 0x00000001ff010000, + 0x00000001ff010001, + 0x00000001ff0101ff, + 0x00000001ff010100, + 0x0000000100ffff00, + 0x0000000100ff0000, + 0x0000000100ff0001, + 0x0000000100ff01ff, + 0x0000000100ff0100, + 0x0000000100ff0101, + 0x000000010000ffff, + 0x000000010000ff00, + 0x000000010000ff01, + 0x00000001000000ff, + 0x0000000100000000, + 0x0000000100000001, + 0x00000001000001ff, + 0x0000000100000100, + 0x0000000100000101, + 0x000000010001ff00, + 0x00000001000100ff, + 0x0000000100010000, + 0x0000000100010100, + 0x0000000101ffff01, + 0x0000000101ff0000, + 0x0000000101ff0001, + 0x0000000101ff01ff, + 0x0000000101ff0100, + 0x0000000101ff0101, + 0x000000010100ff00, + 0x0000000101000000, + 0x0000000101000101, + 0x000000010101ff01, + 0x0000000101010000, + 0x0000000101010001, + 0x00000001010101ff, + 0x0000000101010100, + 0x000001ffffff00ff, + 0x000001ffffff0000, + 0x000001ffffff0001, + 0x000001ffffff0100, + 0x000001ffff00ffff, + 0x000001ffff000000, + 0x000001ffff0001ff, + 0x000001ffff01ff00, + 0x000001ffff010101, + 0x000001ff00ff0000, + 0x000001ff00ff01ff, + 0x000001ff00ff0101, + 0x000001ff0000ff00, + 0x000001ff000000ff, + 0x000001ff00000000, + 0x000001ff00000001, + 0x000001ff000001ff, + 0x000001ff00000100, + 0x000001ff0001ffff, + 0x000001ff0001ff01, + 0x000001ff000100ff, + 0x000001ff00010000, + 0x000001ff01ffff01, + 0x000001ff01ff0100, + 0x000001ff0100ffff, + 0x000001ff0100ff01, + 0x000001ff01000000, + 0x000001ff010001ff, + 0x000001ff0101ff00, + 0x000001ff01010100, + 0x00000100ffffff00, + 0x00000100ffffff01, + 0x00000100ffff0000, + 0x00000100ffff0101, + 0x00000100ff00ff00, + 0x00000100ff0000ff, + 0x00000100ff000000, + 0x00000100ff000001, + 0x00000100ff000100, + 0x00000100ff010000, + 0x0000010000ffff00, + 0x0000010000ff00ff, + 0x0000010000ff0000, + 0x0000010000ff0001, + 0x0000010000ff0100, + 0x000001000000ffff, + 0x000001000000ff00, + 0x000001000000ff01, + 0x00000100000000ff, + 0x0000010000000000, + 0x0000010000000001, + 0x00000100000001ff, + 0x0000010000000100, + 0x0000010000000101, + 0x000001000001ff00, + 0x00000100000100ff, + 0x0000010000010000, + 0x0000010000010001, + 0x0000010000010100, + 0x0000010001ffff00, + 0x0000010001ff0000, + 0x0000010001ff0100, + 0x000001000100ff00, + 0x00000100010000ff, + 0x0000010001000000, + 0x0000010001000001, + 0x00000100010001ff, + 0x0000010001000100, + 0x0000010001010000, + 0x00000101ffff00ff, + 0x00000101ffff01ff, + 0x00000101ff000000, + 0x00000101ff000101, + 0x00000101ff01ffff, + 0x00000101ff010000, + 0x00000101ff010001, + 0x00000101ff010100, + 0x0000010100ff0000, + 0x0000010100ff01ff, + 0x0000010100ff0100, + 0x000001010000ff00, + 0x0000010100000000, + 0x0000010100000001, + 0x00000101000001ff, + 0x0000010100000100, + 0x000001010001ff01, + 0x0000010100010000, + 0x00000101000101ff, + 0x0000010100010101, + 0x0000010101ffff00, + 0x0000010101ff0101, + 0x000001010100ff01, + 0x0000010101000000, + 0x0000010101000001, + 0x00000101010001ff, + 0x0000010101000101, + 0x000001010101ff00, + 0x0001ffffffff0000, + 0x0001ffffff0000ff, + 0x0001ffffff000001, + 0x0001ffffff000100, + 0x0001ffffff010000, + 0x0001ffff00ff00ff, + 0x0001ffff0000ffff, + 0x0001ffff00000000, + 0x0001ffff00000001, + 0x0001ffff000001ff, + 0x0001ffff00000101, + 0x0001ffff0001ff00, + 0x0001ffff000100ff, + 0x0001ffff00010001, + 0x0001ffff00010100, + 0x0001ffff01ffff00, + 0x0001ffff01000001, + 0x0001ffff01010000, + 0x0001ff00ffffff00, + 0x0001ff00ffff00ff, + 0x0001ff00ffff0001, + 0x0001ff00ffff0100, + 0x0001ff00ff00ff01, + 0x0001ff00ff000000, + 0x0001ff00ff01ff00, + 0x0001ff00ff01ff01, + 0x0001ff00ff010001, + 0x0001ff00ff010100, + 0x0001ff0000ff0000, + 0x0001ff0000ff0100, + 0x0001ff000000ff00, + 0x0001ff0000000000, + 0x0001ff0000000001, + 0x0001ff0000000100, + 0x0001ff0000010000, + 0x0001ff0000010001, + 0x0001ff0000010101, + 0x0001ff0001ff00ff, + 0x0001ff0001ff0101, + 0x0001ff000100ff01, + 0x0001ff0001000000, + 0x0001ff000101ff00, + 0x0001ff0001010001, + 0x0001ff0001010100, + 0x0001ff01ff00ff00, + 0x0001ff01ff000001, + 0x0001ff01ff000100, + 0x0001ff0100ffffff, + 0x0001ff0100ffff00, + 0x0001ff0100ff0001, + 0x0001ff0100000000, + 0x0001ff0100000001, + 0x0001ff01000001ff, + 0x0001ff010001ffff, + 0x0001ff0101ff0000, + 0x0001ff010100ff00, + 0x0001ff0101000001, + 0x0001ff0101010000, + 0x000100ffff00ff00, + 0x000100ffff00ff01, + 0x000100ffff000000, + 0x000100ffff000001, + 0x000100ffff000101, + 0x000100ffff01ff00, + 0x000100ffff010001, + 0x000100ffff010100, + 0x000100ff00ffffff, + 0x000100ff00ffff01, + 0x000100ff00ff0000, + 0x000100ff00ff01ff, + 0x000100ff00ff0101, + 0x000100ff0000ff00, + 0x000100ff000000ff, + 0x000100ff00000000, + 0x000100ff00000001, + 0x000100ff00000100, + 0x000100ff00000101, + 0x000100ff0001ffff, + 0x000100ff0001ff01, + 0x000100ff00010000, + 0x000100ff01ff00ff, + 0x000100ff01ff0000, + 0x000100ff01ff0100, + 0x000100ff0100ffff, + 0x000100ff0100ff01, + 0x000100ff010000ff, + 0x000100ff01000000, + 0x000100ff01000001, + 0x000100ff010001ff, + 0x000100ff01000101, + 0x000100ff0101ff00, + 0x000100ff010100ff, + 0x000100ff01010100, + 0x00010000ffff0000, + 0x00010000ffff01ff, + 0x00010000ffff0101, + 0x00010000ff00ff00, + 0x00010000ff000000, + 0x00010000ff000001, + 0x00010000ff000100, + 0x0001000000ff00ff, + 0x0001000000ff0000, + 0x0001000000ff0001, + 0x0001000000ff0100, + 0x000100000000ffff, + 0x000100000000ff00, + 0x00010000000000ff, + 0x0001000000000000, + 0x0001000000000001, + 0x0001000000000100, + 0x000100000001ff00, + 0x00010000000100ff, + 0x0001000000010000, + 0x0001000000010001, + 0x0001000000010100, + 0x0001000001ff0001, + 0x0001000001ff0100, + 0x0001000001ff0101, + 0x000100000100ff00, + 0x0001000001000000, + 0x0001000001000001, + 0x0001000001000100, + 0x0001000001000101, + 0x000100000101ff01, + 0x0001000001010000, + 0x0001000001010001, + 0x00010000010101ff, + 0x00010001ffffff01, + 0x00010001ffff0100, + 0x00010001ff000000, + 0x00010001ff01ffff, + 0x00010001ff010001, + 0x00010001ff0101ff, + 0x00010001ff010100, + 0x0001000100ffffff, + 0x0001000100ff0000, + 0x0001000100ff01ff, + 0x0001000100ff0101, + 0x000100010000ff00, + 0x00010001000000ff, + 0x0001000100000000, + 0x0001000100000001, + 0x00010001000001ff, + 0x0001000100000101, + 0x000100010001ffff, + 0x0001000100010000, + 0x00010001000101ff, + 0x0001000101ffffff, + 0x0001000101ffff01, + 0x0001000101ff0000, + 0x0001000101ff0101, + 0x00010001010000ff, + 0x0001000101000001, + 0x00010001010001ff, + 0x0001000101000100, + 0x000100010101ffff, + 0x00010001010100ff, + 0x0001000101010001, + 0x0001000101010101, + 0x000101ffff000001, + 0x000101ffff000100, + 0x000101ffff010000, + 0x000101ff00ffff00, + 0x000101ff0000ff01, + 0x000101ff00000000, + 0x000101ff00000101, + 0x000101ff0001ff00, + 0x000101ff00010100, + 0x000101ff01ff0000, + 0x000101ff0100ff00, + 0x000101ff010001ff, + 0x000101ff01010001, + 0x00010100ffffff00, + 0x00010100ffff00ff, + 0x00010100ff00ffff, + 0x00010100ff000000, + 0x00010100ff01ff00, + 0x00010100ff0100ff, + 0x00010100ff010001, + 0x00010100ff010100, + 0x0001010000ffffff, + 0x0001010000ffff00, + 0x0001010000ff0000, + 0x0001010000ff0001, + 0x0001010000ff01ff, + 0x000101000000ff00, + 0x00010100000000ff, + 0x0001010000000000, + 0x0001010000000001, + 0x0001010000000100, + 0x000101000001ffff, + 0x0001010000010000, + 0x0001010000010101, + 0x0001010001ffff01, + 0x0001010001ff00ff, + 0x0001010001ff0101, + 0x0001010001000000, + 0x000101000101ff00, + 0x00010100010100ff, + 0x0001010001010000, + 0x0001010001010100, + 0x00010101ff00ff00, + 0x00010101ff000001, + 0x00010101ff0001ff, + 0x0001010100ffff00, + 0x0001010100ff00ff, + 0x0001010100ff0100, + 0x000101010000ffff, + 0x0001010100000000, + 0x00010101000001ff, + 0x0001010100000101, + 0x00010101000100ff, + 0x0001010100010000, + 0x0001010100010100, + 0x0001010101ff0001, + 0x00010101010000ff, + 0x00010101010001ff, + 0x0001010101000101, + 0x0001010101010001, + 0x01ffffffffffffff, + 0x01ffffffffffff01, + 0x01ffffffffff01ff, + 0x01ffffffffff0101, + 0x01ffffffff01ffff, + 0x01ffffffff01ff01, + 0x01ffffffff0101ff, + 0x01ffffffff010101, + 0x01ffffff00ff0000, + 0x01ffffff0000ffff, + 0x01ffffff0000ff00, + 0x01ffffff000000ff, + 0x01ffffff00000001, + 0x01ffffff00000100, + 0x01ffffff00010000, + 0x01ffffff01ffffff, + 0x01ffffff01ffff01, + 0x01ffffff01ff01ff, + 0x01ffffff01ff0101, + 0x01ffffff01000000, + 0x01ffffff0101ffff, + 0x01ffffff0101ff01, + 0x01ffffff010101ff, + 0x01ffffff01010101, + 0x01ffff00ffff0000, + 0x01ffff00ff00ff00, + 0x01ffff00ff0000ff, + 0x01ffff00ff000001, + 0x01ffff00ff000100, + 0x01ffff00ff010000, + 0x01ffff0000ffff00, + 0x01ffff0000ff00ff, + 0x01ffff0000ff0100, + 0x01ffff000000ffff, + 0x01ffff000000ff01, + 0x01ffff0000000000, + 0x01ffff0000000001, + 0x01ffff00000001ff, + 0x01ffff0000000100, + 0x01ffff00000100ff, + 0x01ffff0000010001, + 0x01ffff0000010100, + 0x01ffff0001ff0000, + 0x01ffff0001ff0100, + 0x01ffff00010000ff, + 0x01ffff0001000001, + 0x01ffff0001000100, + 0x01ffff0001010000, + 0x01ffff01ffffffff, + 0x01ffff01ffffff01, + 0x01ffff01ffff01ff, + 0x01ffff01ffff0101, + 0x01ffff01ff000000, + 0x01ffff01ff01ffff, + 0x01ffff01ff01ff01, + 0x01ffff01ff0101ff, + 0x01ffff01ff010101, + 0x01ffff010000ff00, + 0x01ffff01000000ff, + 0x01ffff0100000100, + 0x01ffff0100010000, + 0x01ffff0101ffffff, + 0x01ffff0101ffff01, + 0x01ffff0101ff01ff, + 0x01ffff0101ff0101, + 0x01ffff0101000000, + 0x01ffff010101ffff, + 0x01ffff010101ff01, + 0x01ffff01010101ff, + 0x01ffff0101010101, + 0x01ff00ffff0000ff, + 0x01ff00ffff000100, + 0x01ff00ff00ffff00, + 0x01ff00ff00ff00ff, + 0x01ff00ff0000ff00, + 0x01ff00ff00000000, + 0x01ff00ff00000101, + 0x01ff00ff0001ff00, + 0x01ff00ff000100ff, + 0x01ff00ff00010100, + 0x01ff00ff010000ff, + 0x01ff00ff01000100, + 0x01ff0000ffffff00, + 0x01ff0000ffff0100, + 0x01ff0000ff00ff01, + 0x01ff0000ff000000, + 0x01ff0000ff000101, + 0x01ff0000ff010001, + 0x01ff0000ff010100, + 0x01ff000000ffffff, + 0x01ff000000ffff00, + 0x01ff000000ff0000, + 0x01ff000000ff01ff, + 0x01ff00000000ff00, + 0x01ff0000000000ff, + 0x01ff000000000000, + 0x01ff000000000001, + 0x01ff000000000100, + 0x01ff000000000101, + 0x01ff000000010000, + 0x01ff000000010001, + 0x01ff0000000101ff, + 0x01ff000000010101, + 0x01ff000001ffff00, + 0x01ff000001ff00ff, + 0x01ff000001ff0001, + 0x01ff000001ff0100, + 0x01ff00000100ffff, + 0x01ff00000100ff01, + 0x01ff000001000000, + 0x01ff0000010001ff, + 0x01ff000001010001, + 0x01ff0001ff00ff00, + 0x01ff0001ff000001, + 0x01ff0001ff000100, + 0x01ff0001ff010000, + 0x01ff000100ffff00, + 0x01ff000100ff00ff, + 0x01ff000100ff0100, + 0x01ff000100ff0101, + 0x01ff00010000ffff, + 0x01ff000100000000, + 0x01ff000100000100, + 0x01ff000100000101, + 0x01ff00010001ff00, + 0x01ff000100010001, + 0x01ff000100010101, + 0x01ff000101ff0000, + 0x01ff00010100ff00, + 0x01ff000101000101, + 0x01ff0001010100ff, + 0x01ff01ffffffffff, + 0x01ff01ffffffff01, + 0x01ff01ffffff01ff, + 0x01ff01ffffff0101, + 0x01ff01ffff000000, + 0x01ff01ffff01ffff, + 0x01ff01ffff01ff01, + 0x01ff01ffff0101ff, + 0x01ff01ffff010101, + 0x01ff01ff00ffff00, + 0x01ff01ff00ff0000, + 0x01ff01ff0000ff00, + 0x01ff01ff000000ff, + 0x01ff01ff00000100, + 0x01ff01ff00010000, + 0x01ff01ff00010100, + 0x01ff01ff01ffffff, + 0x01ff01ff01ffff01, + 0x01ff01ff01ff01ff, + 0x01ff01ff01ff0101, + 0x01ff01ff01000000, + 0x01ff01ff0101ffff, + 0x01ff01ff0101ff01, + 0x01ff01ff010101ff, + 0x01ff01ff01010101, + 0x01ff0100ffff0000, + 0x01ff0100ffff0001, + 0x01ff0100ff00ff00, + 0x01ff0100ff0000ff, + 0x01ff0100ff000001, + 0x01ff0100ff010000, + 0x01ff010000ffff00, + 0x01ff010000ff00ff, + 0x01ff010000ff0001, + 0x01ff010000ff0100, + 0x01ff01000000ffff, + 0x01ff01000000ff01, + 0x01ff010000000000, + 0x01ff010000000101, + 0x01ff01000001ff00, + 0x01ff0100000100ff, + 0x01ff010001ff0000, + 0x01ff010001000001, + 0x01ff010001000100, + 0x01ff010001010000, + 0x01ff0101ffffffff, + 0x01ff0101ffffff01, + 0x01ff0101ffff01ff, + 0x01ff0101ffff0101, + 0x01ff0101ff000000, + 0x01ff0101ff01ffff, + 0x01ff0101ff01ff01, + 0x01ff0101ff0101ff, + 0x01ff0101ff010101, + 0x01ff010100ff0000, + 0x01ff01010000ff00, + 0x01ff0101000000ff, + 0x01ff010100000001, + 0x01ff010101ffffff, + 0x01ff010101ffff01, + 0x01ff010101ff01ff, + 0x01ff010101ff0101, + 0x01ff010101000000, + 0x01ff01010101ffff, + 0x01ff01010101ff01, + 0x01ff0101010101ff, + 0x01ff010101010101, + 0x0100ffffffff0000, + 0x0100ffffff00ff00, + 0x0100ffffff000001, + 0x0100ffffff0001ff, + 0x0100ffffff000100, + 0x0100ffffff010000, + 0x0100ffff00ffff00, + 0x0100ffff00ff0001, + 0x0100ffff00ff0100, + 0x0100ffff00000000, + 0x0100ffff000001ff, + 0x0100ffff00000101, + 0x0100ffff00010100, + 0x0100ffff00010101, + 0x0100ffff01ff0000, + 0x0100ffff0100ff00, + 0x0100ffff010000ff, + 0x0100ffff01000001, + 0x0100ffff01000100, + 0x0100ffff01010000, + 0x0100ff00ffffff00, + 0x0100ff00ffff00ff, + 0x0100ff00ffff0001, + 0x0100ff00ffff0100, + 0x0100ff00ff00ffff, + 0x0100ff00ff000000, + 0x0100ff00ff0001ff, + 0x0100ff00ff000101, + 0x0100ff00ff01ff00, + 0x0100ff00ff0100ff, + 0x0100ff00ff010001, + 0x0100ff00ff010100, + 0x0100ff0000ffffff, + 0x0100ff0000ff0000, + 0x0100ff000000ffff, + 0x0100ff000000ff00, + 0x0100ff00000000ff, + 0x0100ff0000000000, + 0x0100ff0000000001, + 0x0100ff0000000100, + 0x0100ff000001ff01, + 0x0100ff0000010000, + 0x0100ff0001ff00ff, + 0x0100ff0001ff0001, + 0x0100ff000100ff01, + 0x0100ff0001000000, + 0x0100ff00010001ff, + 0x0100ff000101ff00, + 0x0100ff00010100ff, + 0x0100ff0001010001, + 0x0100ff0001010100, + 0x0100ff01ffff0000, + 0x0100ff01ff00ff00, + 0x0100ff01ff0000ff, + 0x0100ff01ff000100, + 0x0100ff01ff010000, + 0x0100ff0100ff00ff, + 0x0100ff0100ff0001, + 0x0100ff0100ff0100, + 0x0100ff010000ffff, + 0x0100ff010000ff01, + 0x0100ff0100000000, + 0x0100ff01000001ff, + 0x0100ff0100010001, + 0x0100ff0100010100, + 0x0100ff0101ff0000, + 0x0100ff01010000ff, + 0x0100ff0101000001, + 0x0100ff0101010100, + 0x010000ffffffff00, + 0x010000ffffff00ff, + 0x010000ffffff0001, + 0x010000ffff00ffff, + 0x010000ffff000000, + 0x010000ffff0001ff, + 0x010000ffff010001, + 0x010000ff00ffffff, + 0x010000ff00ff0101, + 0x010000ff0000ff00, + 0x010000ff000000ff, + 0x010000ff00000000, + 0x010000ff00000001, + 0x010000ff000001ff, + 0x010000ff00000100, + 0x010000ff0001ffff, + 0x010000ff0001ff00, + 0x010000ff0001ff01, + 0x010000ff00010000, + 0x010000ff01ff00ff, + 0x010000ff01ff0001, + 0x010000ff0100ff01, + 0x010000ff010000ff, + 0x010000ff01000000, + 0x010000ff010001ff, + 0x010000ff0101ff00, + 0x010000ff01010100, + 0x01000000ffffffff, + 0x01000000ffff0000, + 0x01000000ffff01ff, + 0x01000000ffff0101, + 0x01000000ff00ffff, + 0x01000000ff00ff00, + 0x01000000ff0000ff, + 0x01000000ff000000, + 0x01000000ff000001, + 0x01000000ff000100, + 0x01000000ff01ff00, + 0x01000000ff010000, + 0x01000000ff010100, + 0x01000000ff010101, + 0x0100000000ffff00, + 0x0100000000ff00ff, + 0x0100000000ff0000, + 0x0100000000ff0001, + 0x0100000000ff0100, + 0x010000000000ffff, + 0x010000000000ff00, + 0x010000000000ff01, + 0x01000000000000ff, + 0x0100000000000000, + 0x0100000000000001, + 0x01000000000001ff, + 0x0100000000000100, + 0x0100000000000101, + 0x010000000001ff00, + 0x01000000000100ff, + 0x0100000000010000, + 0x0100000000010001, + 0x0100000000010100, + 0x0100000001ffff00, + 0x0100000001ff0000, + 0x0100000001ff01ff, + 0x010000000100ff00, + 0x010000000100ff01, + 0x01000000010000ff, + 0x0100000001000000, + 0x0100000001000001, + 0x0100000001000100, + 0x0100000001000101, + 0x010000000101ffff, + 0x010000000101ff01, + 0x0100000001010000, + 0x01000000010101ff, + 0x0100000001010101, + 0x01000001ffffff00, + 0x01000001ffff00ff, + 0x01000001ff00ffff, + 0x01000001ff000000, + 0x01000001ff000100, + 0x01000001ff01ffff, + 0x01000001ff010001, + 0x01000001ff010100, + 0x0100000100ff0000, + 0x0100000100ff01ff, + 0x0100000100ff0100, + 0x010000010000ff00, + 0x010000010000ff01, + 0x0100000100000000, + 0x0100000100000001, + 0x0100000100000100, + 0x0100000100010000, + 0x01000001000101ff, + 0x0100000101ffff01, + 0x0100000101ff00ff, + 0x0100000101ff0100, + 0x0100000101ff0101, + 0x010000010100ff01, + 0x01000001010000ff, + 0x0100000101000000, + 0x01000001010100ff, + 0x0100000101010001, + 0x0100000101010100, + 0x010001ffffff0000, + 0x010001ffff000001, + 0x010001ffff000100, + 0x010001ffff010000, + 0x010001ff00ffff00, + 0x010001ff00ff0001, + 0x010001ff0000ffff, + 0x010001ff0000ff01, + 0x010001ff00000000, + 0x010001ff00000001, + 0x010001ff00000101, + 0x010001ff000100ff, + 0x010001ff00010000, + 0x010001ff01ff0000, + 0x010001ff0100ff00, + 0x010001ff01000001, + 0x010001ff01000100, + 0x010001ff01010000, + 0x01000100ffff00ff, + 0x01000100ffff0001, + 0x01000100ffff0100, + 0x01000100ff00ffff, + 0x01000100ff00ff01, + 0x01000100ff000000, + 0x01000100ff0001ff, + 0x01000100ff000101, + 0x01000100ff01ffff, + 0x01000100ff01ff00, + 0x01000100ff0100ff, + 0x01000100ff010001, + 0x0100010000ffffff, + 0x0100010000ffff01, + 0x0100010000ff0000, + 0x0100010000ff01ff, + 0x0100010000ff0101, + 0x010001000000ff00, + 0x01000100000000ff, + 0x0100010000000000, + 0x0100010000000001, + 0x0100010000000100, + 0x010001000001ff01, + 0x0100010000010000, + 0x0100010000010001, + 0x0100010000010101, + 0x0100010001ffff00, + 0x0100010001ff00ff, + 0x010001000100ffff, + 0x010001000100ff01, + 0x0100010001000000, + 0x0100010001000101, + 0x010001000101ff00, + 0x0100010001010001, + 0x01000101ffff0000, + 0x01000101ff000000, + 0x01000101ff010000, + 0x0100010100ff00ff, + 0x0100010100ff0001, + 0x0100010100ff0100, + 0x010001010000ffff, + 0x0100010100000000, + 0x01000101000001ff, + 0x010001010001ff00, + 0x0100010101ff0000, + 0x010001010100ff00, + 0x01000101010000ff, + 0x0100010101000000, + 0x0100010101000001, + 0x0101ffffffffffff, + 0x0101ffffffffff01, + 0x0101ffffffff01ff, + 0x0101ffffffff0101, + 0x0101ffffff000000, + 0x0101ffffff01ffff, + 0x0101ffffff01ff01, + 0x0101ffffff0101ff, + 0x0101ffffff010101, + 0x0101ffff00ff0000, + 0x0101ffff0000ff00, + 0x0101ffff000000ff, + 0x0101ffff00000001, + 0x0101ffff00000100, + 0x0101ffff01ffffff, + 0x0101ffff01ffff01, + 0x0101ffff01ff01ff, + 0x0101ffff01ff0101, + 0x0101ffff01000000, + 0x0101ffff0101ffff, + 0x0101ffff0101ff01, + 0x0101ffff010101ff, + 0x0101ffff01010101, + 0x0101ff00ffff0000, + 0x0101ff00ffff0100, + 0x0101ff00ff00ff00, + 0x0101ff00ff0000ff, + 0x0101ff00ff000001, + 0x0101ff00ff000100, + 0x0101ff00ff000101, + 0x0101ff0000ff0001, + 0x0101ff0000ff0100, + 0x0101ff000000ff00, + 0x0101ff0000000000, + 0x0101ff00000001ff, + 0x0101ff0000000101, + 0x0101ff000001ff00, + 0x0101ff00000100ff, + 0x0101ff0001ff0000, + 0x0101ff000100ffff, + 0x0101ff000100ff01, + 0x0101ff0001000001, + 0x0101ff0001000100, + 0x0101ff01ffffff01, + 0x0101ff01ffff01ff, + 0x0101ff01ffff0101, + 0x0101ff01ff00ffff, + 0x0101ff01ff000100, + 0x0101ff01ff01ff01, + 0x0101ff01ff0101ff, + 0x0101ff01ff010101, + 0x0101ff0100ff0000, + 0x0101ff010000ff00, + 0x0101ff0100000001, + 0x0101ff0100000100, + 0x0101ff0100010000, + 0x0101ff0101ffffff, + 0x0101ff0101ffff01, + 0x0101ff0101ff01ff, + 0x0101ff0101ff0101, + 0x0101ff0101000000, + 0x0101ff010101ffff, + 0x0101ff010101ff01, + 0x0101ff01010101ff, + 0x0101ff0101010101, + 0x010100ffff000100, + 0x010100ffff010000, + 0x010100ff00ffff00, + 0x010100ff00ff00ff, + 0x010100ff0000ffff, + 0x010100ff000000ff, + 0x010100ff00000000, + 0x010100ff000001ff, + 0x010100ff00000101, + 0x010100ff0001ff00, + 0x010100ff00010000, + 0x010100ff00010001, + 0x010100ff000101ff, + 0x010100ff00010100, + 0x010100ff01ff0000, + 0x01010000ffff0001, + 0x01010000ffff0100, + 0x01010000ff00ffff, + 0x01010000ff00ff01, + 0x01010000ff000000, + 0x01010000ff0001ff, + 0x01010000ff010001, + 0x01010000ff010100, + 0x0101000000ffff01, + 0x0101000000ff0000, + 0x010100000000ff00, + 0x01010000000000ff, + 0x0101000000000000, + 0x0101000000000001, + 0x0101000000000100, + 0x0101000000010000, + 0x0101000000010101, + 0x0101000001ffff00, + 0x0101000001ff00ff, + 0x0101000001ff0000, + 0x0101000001ff0001, + 0x0101000001ff0100, + 0x010100000100ff01, + 0x0101000001000000, + 0x01010000010001ff, + 0x01010001ffff0000, + 0x01010001ff00ff00, + 0x01010001ff000001, + 0x01010001ff000101, + 0x01010001ff01ff00, + 0x01010001ff010000, + 0x0101000100ff00ff, + 0x0101000100ff0001, + 0x0101000100ff0101, + 0x010100010000ff01, + 0x0101000100000000, + 0x0101000100000001, + 0x01010001000001ff, + 0x010100010001ffff, + 0x010100010001ff01, + 0x0101000101ff0001, + 0x010100010100ffff, + 0x0101000101000000, + 0x0101000101000001, + 0x0101000101000100, + 0x010100010101ff00, + 0x01010001010100ff, + 0x0101000101010001, + 0x010101ffffffffff, + 0x010101ffffffff01, + 0x010101ffffff01ff, + 0x010101ffffff0101, + 0x010101ffff01ffff, + 0x010101ffff01ff01, + 0x010101ffff0101ff, + 0x010101ffff010101, + 0x010101ff0000ff00, + 0x010101ff000000ff, + 0x010101ff00000001, + 0x010101ff00000100, + 0x010101ff01ffffff, + 0x010101ff01ffff01, + 0x010101ff01ff01ff, + 0x010101ff01ff0101, + 0x010101ff01000000, + 0x010101ff0101ffff, + 0x010101ff0101ff01, + 0x010101ff010101ff, + 0x010101ff01010101, + 0x01010100ffff0000, + 0x01010100ff0000ff, + 0x01010100ff000100, + 0x01010100ff01ff00, + 0x01010100ff010000, + 0x0101010000ffff00, + 0x010101000000ffff, + 0x0101010000000000, + 0x0101010000000101, + 0x010101000001ff00, + 0x0101010000010001, + 0x0101010000010100, + 0x010101000100ffff, + 0x0101010001000001, + 0x01010101ffffffff, + 0x01010101ffffff01, + 0x01010101ffff01ff, + 0x01010101ffff0101, + 0x01010101ff01ffff, + 0x01010101ff01ff01, + 0x01010101ff0101ff, + 0x01010101ff010101, + 0x010101010000ff00, + 0x01010101000000ff, + 0x0101010100000001, + 0x0101010101ffffff, + 0x0101010101ffff01, + 0x0101010101ff01ff, + 0x0101010101ff0101, + 0x0101010101000000, + 0x010101010101ffff, + 0x010101010101ff01, + 0x01010101010101ff, + 0x0101010101010101, +}; +} // namespace iq_grids diff --git a/apps/ggml/halide/k_quant_generators.cpp b/apps/ggml/halide/k_quant_generators.cpp new file mode 100644 index 000000000000..319c95de8ba8 --- /dev/null +++ b/apps/ggml/halide/k_quant_generators.cpp @@ -0,0 +1,80 @@ +// Generic, GeneratorParam-driven quantize/dequantize pair for GGML's +// K-quant super-block formats (see quant_components.h's make_k_quant_scheme/ +// CombineBits/PlanarBitPack/K4ScaleMinPack/Q3KScalePack for the reusable +// Approximation pieces this assembles). "Q2_K"/"Q3_K"/"Q4_K"/"Q5_K"/"Q6_K" are not +// distinct C++ classes here -- they're just different GENERATOR_ARGS +// instantiations of the same generator template, registered in +// CMakeLists.txt as q4_k_quantize/q4_k_dequantize etc. -- following the +// exact same shape as lookup_table_quant_generators.cpp's +// LookupTableCodecGenerator (see that file's header comment for +// the full rationale: one shared configure() builds the whole pipeline once +// from a real ImageParam, each direction just adopts whichever half applies +// via add_input(const ImageParam&)/add_output(const Func&); generate() is +// an empty stub). +// +// generate() never calls Approximation::encode()/decode() directly -- only +// through Func::approximate_by() and Pipeline::sever(). This +// configure()/generate() body is identical across every *_quant_generators.cpp +// file in this directory, so it lives in codec_generator_base.h's +// CodecGeneratorBase instead of being repeated here -- this +// class only needs to supply its own GeneratorParams and a build_scheme(). + +#include "Halide.h" + +#include "codec_generator_base.h" +#include "quant_components.h" + +using namespace Halide; +using namespace ggml_halide; + +namespace { + +// Which of quant_components.h's make_*_scheme() factories to use, and the +// on-disk block size that goes with it -- not derivable from a plain +// GeneratorParam combination, since each K-quant format's field layout, +// scale scheme, and code bit-width all differ. +enum class Family { Q2_K, + Q3_K, + Q4_K, + Q5_K, + Q6_K }; + +template +class KQuantCodecGenerator : public CodecGeneratorBase, dir> { +public: + GeneratorParam family{ + "family", + Family::Q4_K, + {{"q2_k", Family::Q2_K}, + {"q3_k", Family::Q3_K}, + {"q4_k", Family::Q4_K}, + {"q5_k", Family::Q5_K}, + {"q6_k", Family::Q6_K}}}; + + SchemeAndBytes build_scheme() const { + // switch's controlling expression can't resolve GeneratorParam's + // implicit conversion operators unambiguously -- .value() sidesteps + // that by returning the plain Family directly. Each make_*_scheme() + // now returns its own block_bytes alongside the scheme (computed from + // the same field list it builds internally), so there's no byte + // arithmetic to duplicate here. + switch (family.value()) { + case Family::Q2_K: + return make_q2_k_scheme(); // {scales[16]; qs[64]; fp16 d; fp16 dmin;} + case Family::Q3_K: + return make_q3_k_scheme(); // {hmask[32]; qs[64]; scales[12]; fp16 d;} + case Family::Q4_K: + return make_q4_k_scheme(); // {fp16 d; fp16 dmin; scales[12]; qs[128];} + case Family::Q5_K: + return make_q5_k_scheme(); // {fp16 d; fp16 dmin; scales[12]; qh[32]; qs[128];} + case Family::Q6_K: + return make_q6_k_scheme(); // {ql[128]; qh[64]; scales[16]; fp16 d;} + } + _halide_internal_error << "unreachable Family\n"; + } +}; + +} // namespace + +HALIDE_REGISTER_GENERATOR(KQuantCodecGenerator, k_quant_quantize) +HALIDE_REGISTER_GENERATOR(KQuantCodecGenerator, k_quant_dequantize) diff --git a/apps/ggml/halide/k_quant_vec_dot_generator.cpp b/apps/ggml/halide/k_quant_vec_dot_generator.cpp new file mode 100644 index 000000000000..0468970747f9 --- /dev/null +++ b/apps/ggml/halide/k_quant_vec_dot_generator.cpp @@ -0,0 +1,63 @@ +// Generic, family-driven vec_dot for the K-quant super-block formats, the +// vec_dot counterpart of k_quant_generators.cpp's KQuantCodecGenerator (same +// family set). "q4_k_vec_dot" is a PARAMS family=q4_k instantiation of this +// one generator. Weight is a block-indexed K-quant codec, activation is the +// block-indexed Q8_K codec; VecDotGeneratorBase splices both via +// approximate_by/sever. K-quant decode is a two-level (sub-block) +// scale, so the per-block scale is not single-invariant -> Float schedule. + +#include "Halide.h" + +#include "quant_components.h" +#include "vec_dot_generator_base.h" + +using namespace Halide; +using namespace ggml_halide; + +namespace { + +enum class Family { Q2_K, + Q3_K, + Q4_K, + Q5_K, + Q6_K }; + +class KQuantVecDotGenerator : public VecDotGeneratorBase { +public: + GeneratorParam family{ + "family", + Family::Q4_K, + {{"q2_k", Family::Q2_K}, + {"q3_k", Family::Q3_K}, + {"q4_k", Family::Q4_K}, + {"q5_k", Family::Q5_K}, + {"q6_k", Family::Q6_K}}}; + + // Q8_K activation codec (block_q8_K = {float d; qs[256]; bsums[16]} = 292 bytes). + static Halide::Approximation q8_k_codec() { + return make_q8_k_scheme(256, 127, Layout::BlockIndexed).scheme; + } + + VecDotSpec build_vec_dot() const { + // All K-quants: 256-element super-block, Q8_K activation, two-level + // scale -> Float schedule. + switch (family.value()) { + case Family::Q2_K: + return {make_q2_k_scheme(Layout::BlockIndexed).scheme, 84, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::Q3_K: + return {make_q3_k_scheme(Layout::BlockIndexed).scheme, 110, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::Q4_K: + return {make_q4_k_scheme(Layout::BlockIndexed).scheme, 144, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::Q5_K: + return {make_q5_k_scheme(Layout::BlockIndexed).scheme, 176, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::Q6_K: + return {make_q6_k_scheme(Layout::BlockIndexed).scheme, 210, q8_k_codec(), 292, 256, ScheduleKind::Float}; + } + _halide_internal_error << "KQuantVecDotGenerator: family not yet converted\n"; + return {}; + } +}; + +} // namespace + +HALIDE_REGISTER_GENERATOR(KQuantVecDotGenerator, k_quant_vec_dot) diff --git a/apps/ggml/halide/lookup_table_quant_generators.cpp b/apps/ggml/halide/lookup_table_quant_generators.cpp new file mode 100644 index 000000000000..aa284c0a0263 --- /dev/null +++ b/apps/ggml/halide/lookup_table_quant_generators.cpp @@ -0,0 +1,116 @@ +// Generic, GeneratorParam-driven quantize/dequantize pair for GGML's +// codebook-quantized (lookup-table) formats (see quant_components.h's +// LookupTableQuantize/E8M0Pack for the reusable Approximation pieces this +// assembles). "IQ4_NL"/"MXFP4" are not distinct C++ classes here -- they're +// just different GENERATOR_ARGS instantiations of the same generator +// template, registered in CMakeLists.txt as iq4_nl_quantize/mxfp4_quantize +// etc. -- following the same shape as symmetric_quant_generators.cpp's +// SymmetricCodecGenerator (see that file's header comment for the +// full rationale: one shared configure() builds the whole pipeline once from +// a real ImageParam, each direction just adopts whichever half applies via +// add_input(const ImageParam&)/add_output(const Func&); generate() is an +// empty stub). +// +// generate() never calls Approximation::encode()/decode() directly -- only +// through Func::approximate_by() and Pipeline::sever(). This +// configure()/generate() body is identical across every *_quant_generators.cpp +// file in this directory, so it lives in codec_generator_base.h's +// CodecGeneratorBase instead of being repeated here -- this +// class only needs to supply its own GeneratorParams and a build_scheme(). + +#include "Halide.h" + +#include "codec_generator_base.h" +#include "quant_components.h" + +using namespace Halide; +using namespace ggml_halide; + +namespace { + +// Which of quant_components.h's make_*_scheme() factories to use, and the +// on-disk block size that goes with it -- not derivable from a plain +// GeneratorParam combination the way the affine/symmetric family's +// SchemeKind's block_bytes is, since each codebook has its own fixed layout. +enum class Family { IQ4_NL, + MXFP4, + TQ2_0, + TQ1_0, + NVFP4, + IQ2_S, + IQ3_XXS, + IQ3_S, + IQ4_XS, + IQ2_XS, + IQ2_XXS, + IQ1_S, + IQ1_M }; + +template +class LookupTableCodecGenerator : public CodecGeneratorBase, dir> { +public: + GeneratorParam family{ + "family", + Family::IQ4_NL, + {{"iq4_nl", Family::IQ4_NL}, + {"mxfp4", Family::MXFP4}, + {"tq2_0", Family::TQ2_0}, + {"tq1_0", Family::TQ1_0}, + {"nvfp4", Family::NVFP4}, + {"iq2_s", Family::IQ2_S}, + {"iq3_xxs", Family::IQ3_XXS}, + {"iq3_s", Family::IQ3_S}, + {"iq4_xs", Family::IQ4_XS}, + {"iq2_xs", Family::IQ2_XS}, + {"iq2_xxs", Family::IQ2_XXS}, + {"iq1_s", Family::IQ1_S}, + {"iq1_m", Family::IQ1_M}}}; + + SchemeAndBytes build_scheme() const { + // switch's controlling expression can't resolve GeneratorParam's + // implicit conversion operators unambiguously -- .value() sidesteps + // that by returning the plain Family directly. Every make_*_scheme() + // now returns its own block_bytes alongside the scheme (from its + // field table, or -- for the IQ grid leaves, which are deliberately + // NOT field-table-decomposed; see quant_components.h section 6's + // design note -- the hand-verified constant declared next to the + // leaf), so there's no literal byte count to keep in sync here. + switch (family.value()) { + case Family::IQ4_NL: + return make_iq4_nl_scheme(); + case Family::MXFP4: + return make_mxfp4_scheme(); + case Family::TQ2_0: + return make_tq2_0_scheme(); + case Family::TQ1_0: + return make_tq1_0_scheme(); + case Family::NVFP4: + return make_nvfp4_scheme(); + case Family::IQ2_S: + return make_iq2_s_scheme(); // {fp16 d; qs[32]; signs[32]; qh[8]; scales[8];} + case Family::IQ3_XXS: + return make_iq3_xxs_scheme(); // {fp16 d; qs[64]; scales_and_signs[32];} + case Family::IQ3_S: + return make_iq3_s_scheme(); // {fp16 d; qs[64]; qh[8]; signs[32]; scales[4];} + case Family::IQ4_XS: + return make_iq4_xs_scheme(); // {fp16 d; scales_h[2]; scales_l[4]; qs[128];} + // Importance-matrix-only (dequantize direction only -- no quantize + // library is built for these; SeveredEncode stands in for the missing + // forward map). + case Family::IQ2_XS: + return make_iq2_xs_scheme(); + case Family::IQ2_XXS: + return make_iq2_xxs_scheme(); + case Family::IQ1_S: + return make_iq1_s_scheme(); + case Family::IQ1_M: + return make_iq1_m_scheme(); + } + _halide_internal_error << "unreachable Family\n"; + } +}; + +} // namespace + +HALIDE_REGISTER_GENERATOR(LookupTableCodecGenerator, lookup_table_quantize) +HALIDE_REGISTER_GENERATOR(LookupTableCodecGenerator, lookup_table_dequantize) diff --git a/apps/ggml/halide/lookup_table_vec_dot_generator.cpp b/apps/ggml/halide/lookup_table_vec_dot_generator.cpp new file mode 100644 index 000000000000..9500d613b507 --- /dev/null +++ b/apps/ggml/halide/lookup_table_vec_dot_generator.cpp @@ -0,0 +1,111 @@ +// Generic, family-driven vec_dot for the codebook/grid formats, the vec_dot +// counterpart of lookup_table_quant_generators.cpp's LookupTableCodecGenerator +// (same family set). "iq4_nl_vec_dot" is not a distinct C++ class -- it's a +// PARAMS family=iq4_nl instantiation of this one generator, registered in +// CMakeLists.txt. Weight and activation are both block-indexed codecs from +// quant_components.h; VecDotGeneratorBase splices them via approximate_by/ +// sever (see vec_dot_generator_base.h). + +#include "Halide.h" + +#include "quant_components.h" +#include "vec_dot_generator_base.h" + +using namespace Halide; +using namespace ggml_halide; + +namespace { + +enum class Family { IQ4_NL, + MXFP4, + TQ2_0, + TQ1_0, + NVFP4, + IQ2_S, + IQ3_XXS, + IQ3_S, + IQ4_XS, + IQ2_XS, + IQ2_XXS, + IQ1_S, + IQ1_M }; + +class LookupTableVecDotGenerator : public VecDotGeneratorBase { +public: + GeneratorParam family{ + "family", + Family::IQ4_NL, + {{"iq4_nl", Family::IQ4_NL}, + {"mxfp4", Family::MXFP4}, + {"tq2_0", Family::TQ2_0}, + {"tq1_0", Family::TQ1_0}, + {"nvfp4", Family::NVFP4}, + {"iq2_s", Family::IQ2_S}, + {"iq3_xxs", Family::IQ3_XXS}, + {"iq3_s", Family::IQ3_S}, + {"iq4_xs", Family::IQ4_XS}, + {"iq2_xs", Family::IQ2_XS}, + {"iq2_xxs", Family::IQ2_XXS}, + {"iq1_s", Family::IQ1_S}, + {"iq1_m", Family::IQ1_M}}}; + + // Q8_0 activation codec (block_q8_0 = {fp16 d; qs[32]} = 34 bytes). + static Halide::Approximation q8_0_codec() { + return make_symmetric_block_scheme(32, 127, RoundingMode::Nearest, ScaleAnchor::AbsMax, 8, Layout::BlockIndexed).scheme; + } + // Q8_K activation codec (block_q8_K = {float d; qs[256]; bsums[16]} = 292 bytes). + static Halide::Approximation q8_k_codec() { + return make_q8_k_scheme(256, 127, Layout::BlockIndexed).scheme; + } + + VecDotSpec build_vec_dot() const { + switch (family.value()) { + case Family::IQ4_NL: + // 4-bit codebook, single fp16 scale x Q8_0: single per-block scale + // and int8 codebook values -> SDOT-eligible (the base header's + // deep-inline SDOT exposes the scale even through the codebook LUT). + return {make_iq4_nl_scheme(Layout::BlockIndexed).scheme, 18, q8_0_codec(), 34, 32, ScheduleKind::SDOT}; + case Family::MXFP4: + // Same single-scale codebook shape as IQ4_NL (E8M0 scale) x Q8_0. + return {make_mxfp4_scheme(Layout::BlockIndexed).scheme, 17, q8_0_codec(), 34, 32, ScheduleKind::SDOT}; + case Family::NVFP4: + // 64-element block (4 sub-scales) x Q8_0 (32-block): the activation + // is Reblocked 32 -> 64 so both share the weight's block. Sub-block + // scales -> Float. + return {make_nvfp4_scheme(Layout::BlockIndexed).scheme, 36, reblock_activation(q8_0_codec(), 32, 64), 34, 64, ScheduleKind::Float}; + // TQ1_0/TQ2_0 x Q8_K, IQ4_XS x Q8_K: single fp16 scale (TQ) or two-level + // scale (IQ4_XS) -> Float schedule for now. + case Family::TQ2_0: + return {make_tq2_0_scheme(Layout::BlockIndexed).scheme, 66, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::TQ1_0: + return {make_tq1_0_scheme(Layout::BlockIndexed).scheme, 54, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::IQ4_XS: + return {make_iq4_xs_scheme(Layout::BlockIndexed).scheme, 136, q8_k_codec(), 292, 256, ScheduleKind::Float}; + // Grid formats x Q8_K: per-group scale + sign -> Float schedule. The + // block-indexed codec collapses the leaf's {8,4,8} output to (kk, blk). + case Family::IQ2_S: + return {make_iq2_s_scheme(Layout::BlockIndexed).scheme, 82, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::IQ3_XXS: + return {make_iq3_xxs_scheme(Layout::BlockIndexed).scheme, 98, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::IQ3_S: + return {make_iq3_s_scheme(Layout::BlockIndexed).scheme, 110, q8_k_codec(), 292, 256, ScheduleKind::Float}; + // Importance-matrix-only formats (SeveredEncode weight scheme) x Q8_K. + case Family::IQ2_XS: + return {make_iq2_xs_scheme(Layout::BlockIndexed).scheme, 74, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::IQ2_XXS: + return {make_iq2_xxs_scheme(Layout::BlockIndexed).scheme, 66, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::IQ1_S: + return {make_iq1_s_scheme(Layout::BlockIndexed).scheme, 50, q8_k_codec(), 292, 256, ScheduleKind::Float}; + case Family::IQ1_M: + return {make_iq1_m_scheme(Layout::BlockIndexed).scheme, 56, q8_k_codec(), 292, 256, ScheduleKind::Float}; + default: + break; + } + _halide_internal_error << "LookupTableVecDotGenerator: family not yet converted\n"; + return {}; + } +}; + +} // namespace + +HALIDE_REGISTER_GENERATOR(LookupTableVecDotGenerator, lookup_table_vec_dot) diff --git a/apps/ggml/halide/quant_components.h b/apps/ggml/halide/quant_components.h new file mode 100644 index 000000000000..e466dd13a7e2 --- /dev/null +++ b/apps/ggml/halide/quant_components.h @@ -0,0 +1,3521 @@ +#pragma once + +// Reusable Approximation components for GGML-style per-block quantized +// weight formats -- see doc/Approximation.md for the rationale. Every weight format is built by composing +// these kinds of pieces via Halide::Compose/Halide::Parallel (and, for the +// extern-delegated formats, Halide::TrustedInverse) into a scheme (see the +// make_*_scheme() factory functions below), which the +// Generators then splice in via Func::approximate_by()/ +// Pipeline::sever() -- never by calling Approximation::encode()/ +// decode() directly. +// +// 1. BlockReshape -- lossless relayout: flat values <-> (kk, blk). +// 2. SymmetricAffineQuantize/AffineQuantize -- the actual lossy step: +// block values <-> (integer codes, one or two float(s) per block). +// 3a. Fp16Pack/PlanarBitPack/BytePack -- per-field packing: +// a typed field (codes, scale, min) <-> its own on-disk byte encoding. +// 3b. AppendSums -- a derived extra field, computed from other +// already-encoded fields rather than from the original values. +// 3c. StructPack -- concatenates N already-packed fields into one +// byte-addressed buffer, matching a specific on-disk block layout +// (e.g. block_q4_0). +// 4. Extern-delegated formats (codebook/K-quant/IQ grid/IQ4_XS): their +// quantize is an opaque GGML extern (ExternQuantize), paired with a +// compositional dequantize via Halide::TrustedInverse. The extra +// decode-only math leaves those need -- Codebook (codes -> table[codes]), +// LinearDequant (the scale multiply), CombineBits +// (K-quant split codes) -- live in section 4 alongside ExternQuantize. +// +// None of these know about any specific GGML type name -- "Q4_0"/"Q4_1"/etc. +// are just particular parameter choices, assembled where the Generators live +// (symmetric_quant_generators.cpp), not encoded here. +// +// None of these components call Func::bound() on their own intermediate +// Funcs, even where a range is intrinsically known (e.g. codes(kk, blk) for +// kk in [0, block_size)): with everything left at its default (inline) +// schedule, as it is here, Halide already infers the true required range by +// propagating backward from wherever the final packed buffer is actually +// realized or scheduled (a Generator's Output dim bounds, or an explicit +// realize() shape) -- an explicit bound() on an inlined Func is simply +// ignored ("meaningless... because the function is scheduled inline", per +// Halide's own warning). bound() only becomes useful once a component's Func +// is deliberately scheduled non-inline (compute_root/compute_at), which is a +// scheduling-time decision made where that happens, not decided here. +// +// DIMENSION / WILDCARD CONVENTION: a component's Func indices are laid out as +// (field dims..., blk, lane dims...) -- the field's own within-block dims +// first (kk, or byte, or (plane, sub), ...), then the block index blk, then +// any trailing "lane" dims (e.g. a repack matmul weight's column-in-group j +// and col-group x), carried by the Halide::_ placeholder. Decode-only stages +// should be lane-general by default: write decode() with a trailing +// Halide::_ on both sides (f(kk, blk, _) = g(..., blk, _)) so the same +// component runs unchanged whether there are zero lane dims (the plain +// (kk, blk) codecs) or several (the repack weight schemes). Encode stages +// stay at fixed arity -- every encoder here runs on a plain (kk, blk) or +// flat row, so there is nothing for a wildcard to carry. Currently +// lane-general (decode side): Fp16Pack, F32Pack, E8M0Pack, BytePack, +// PlanarBitPack (both normal and plane-axis modes), Codebook, LinearDequant, +// K4ScaleMinPack, Q3KScalePack, IQ4XSScalePack. The remaining decoders +// (BlockReshape/Reblock, the repack interleave/de-interleave leaves, the IQ +// grid leaves, TritPack, BitPack, Int16Pack, UE4M3Pack) are fixed-arity by +// design or have no lane-general consumer yet. The full dimension-general +// BlockLayout sketched in the DESIGN NOTE by Reblock (splitting/permuting +// arbitrary dims) remains deliberately deferred. + +// Only the aggregated Halide.h is installed for apps to consume (individual +// per-class headers like Approximation.h are not) -- it already pulls in +// Approximation/Compose/Parallel/Pipeline::sever. +#include "Halide.h" + +#include +#include +#include +#include + +#include "iq_grids_data.h" + +namespace ggml_halide { + +// The result of a make_*_scheme() factory: the scheme itself plus its +// on-disk block byte count, computed once from the same field list +// make_block_layout() (see FieldSpec below) already sums, rather than +// hand-summed again at each Generator call site. Shared with +// codec_generator_base.h's CodecGeneratorBase (which is what actually +// consumes it -- see there). +struct SchemeAndBytes { + Halide::Approximation scheme; + int block_bytes; + // When set (is_struct()), the scheme's encoded form is a first-class + // Type::Struct block (one struct per block index) rather than a 2-D + // (byte, blk) UInt(8) buffer. The Generator uses this to declare a + // struct-typed, 1-D packed ImageParam/Output, and block_bytes is then just + // block_type.bytes() -- the single source of truth for the on-disk width. + // Default-constructed (invalid) for the byte-buffer schemes not yet ported. + Halide::Type block_type; + // Set when the scheme appends a per-block scaled sum of its codes (Q8_1's + // `s` field -- AppendSums{SumMode::ScaledFloat}). A vec_dot pairing an + // affine weight against such an activation can sever the offset term's + // sum(act) accumulator straight to this stored field instead of recomputing + // it -- see VecDotGeneratorBase::configure(). Purely a byte-path, + // 32-element-block property today (Q8_1); wider layouts leave it false. + bool has_block_sums = false; + // Optional scheduling identities exported by composed schemes. q5_0 uses + // these to materialize reconstructed codes and its packed qh word without + // relying on generated Func names. Legacy q5_1 intentionally leaves them + // undefined (see VecDotGeneratorBase::configure()). + Halide::Approximation reconstructed_codes_stage; + Halide::Approximation packed_high_word_stage; +}; + +// Every make_*_scheme() factory below takes a Layout, selecting what its +// BlockReshape (or grid BlockReshape) does with the "flat" side: +// - FlatRow (the default): a fully-flat 1-D row -- the shape +// quantize_row/dequantize_row Generators want. +// - BlockIndexed: a passthrough (kk, blk) -- the shape a vec_dot/repack +// Generator wants (its own reduction already runs over (kk, blk); no +// flat<->block reshape is needed on top). This used to be a *separate* +// make_*_codec() function per scheme (build the codec, skip the +// reshape); now it's the same factory with BlockReshape's own +// block_indexed flag set true, which makes it a lossless identity +// passthrough (see BlockReshape's own comment) -- so there's no +// behavioral difference, just one fewer named entry point per scheme. +enum class Layout { FlatRow, + BlockIndexed }; + +// --------------------------------------------------------------------------- +// 1. Lossless relayout. +// --------------------------------------------------------------------------- + +// A lossless flat <-> block reshape. In the common one-dimensional case +// (BlockReshape(block_size)), packed(kk, blk) = flat(blk*block_size + kk) -- +// one within-block index kk in [0, block_size), one block index blk. +// +// The general case (BlockReshape({e0, e1, ...})) unflattens the within-block +// index into *several* dimensions, innermost/fastest-varying first, so a +// component whose values have nested block structure can index those +// dimensions directly instead of re-deriving them from a flat kk via div/mod. +// E.g. an IQ 256-element superblock structured as group(8) x l(4) x elem(8) +// uses extents {8, 4, 8}: packed(elem, l, group, blk), where the flat +// within-block index kk = elem + 8*l + 32*group (product of extents = block +// size). The single-int constructor is exactly the one-extent case. +// +// `block_indexed` selects what the "flat" side looks like: +// - false (default): a fully-flat 1-D row f(k), k = blk*block_size + within +// -- the shape quantize_row/dequantize_row want. +// - true: a block-indexed 2-D f(kk, blk), within-block index kept separate +// from the block index -- the shape a per-block vec_dot reduction wants +// (so the block index stays a distinct RVar for the SDOT rfactor hoist). +// In block-indexed mode a single-extent reshape is a (kk,blk) passthrough, +// and a multi-extent one collapses the nested dims (elem,l,group,blk) into +// (kk,blk) -- the only difference from flat mode is folding blk into k or not. +// Compatibility alias: the implementation is now a public core component. +using BlockReshape = Halide::BlockReshape; + +// --------------------------------------------------------------------------- +// Lossless block-layout relayouts (Reblock, and the repack Interleave below). +// +// DESIGN NOTE (intended library form, deferred): the clean, general shape for +// these is a *dimension-general* block-relayout -- a component that splits and +// permutes arbitrary index dimensions, carrying any trailing dims through +// untouched (the way the Python research sketch's BlockLayout(splits=...)/ +// SplitStorage do, via Halide's `_` placeholder), quantizing/packing whatever +// falls out. That would let a single `BlockLayout` utility live in the core +// Approximation library and compose in front of *any* lossy quant, with the +// repack interleave being just one instantiation. We deliberately do NOT do +// that here: `Func`/`Var` `_` wildcard semantics were a source of trouble, and +// pinning them down is its own rabbit hole. Instead these relayouts are +// written at fixed, concrete arities (folding any extra "lane" -- e.g. repack's +// 4 interleaved rows -- into the block index blk), reusing the existing +// (kk, blk) quant/pack components unchanged. When the wildcard story is sorted, +// these should graduate to the dimension-general form. +// --------------------------------------------------------------------------- + +// Losslessly re-view a block-indexed Func at a different block size (the flat +// element order is unchanged; only the (kk, blk) factorization differs). +// decode(): (kk_from, blk_from) at `from_block` -> (kk_to, blk_to) at +// `to_block`, reading the same global element g = blk_to*to_block + kk_to from +// its source position (g % from_block, g / from_block). encode() is the +// mirror. This is what lets a vec_dot present an activation stored in its own +// (smaller) block size -- e.g. Q8_0's 32-element blocks -- at a weight's +// (larger) block size (Q1_0's 128, NVFP4's 64), so both operands share one +// (kk, blk) and the Generator's reduction stays uniform. The block-structure +// reconciliation lives here, in an Approximation, not open-coded in the +// reduction. +class Reblock { +public: + Reblock(int from_block, int to_block) + : from_(from_block), to_(to_block) { + } + + Halide::Func encode(const Halide::Func &in) const { + using namespace Halide; + // (kk, blk) at to_block + Var kk("kk"), blk("blk"); + Expr g = blk * from_ + kk; + Func out("reblock_encoded"); + out(kk, blk) = in(g % to_, g / to_); + return out; + } + + Halide::Func decode(const Halide::Func &in) const { + using namespace Halide; + // (kk, blk) at from_block + Var kk("kk"), blk("blk"); + Expr g = blk * to_ + kk; + Func out("reblock_decoded"); + out(kk, blk) = in(g % from_, g / from_); + return out; + } + +private: + int from_, to_; +}; + +// An activation codec that decodes to (kk, blk) at `to_block`: `act_codec` +// (block-indexed at the activation's own `from_block`) composed with a Reblock +// when the two differ, else `act_codec` unchanged. Lets a vec_dot pair a weight +// of one block size against an activation of another (Q1_0/NVFP4 x Q8_0). +inline Halide::Approximation reblock_activation( + Halide::Approximation act_codec, int from_block, int to_block) { + if (from_block == to_block) { + return act_codec; + } + return Halide::Compose(Reblock{from_block, to_block}, std::move(act_codec)); +} + +// --------------------------------------------------------------------------- +// 2. The lossy step. +// --------------------------------------------------------------------------- + +// GGML's own reference quantizers round differently depending on the target +// bit width, not out of taste but because they use different formulas: +// - Nearest: plain round-half-away-from-zero (Q8_0's quantize_row_q8_0_ref +// uses roundf()). Halide's round() matches this exactly. +// - TruncateHalfUpWithOffset: a truncate-based "+qmax+0.5f then cast" +// trick used by nibble-packed formats (Q4_0's quantize_row_q4_0_ref). +// Verified by hand this is round-half-*up*, not round-half-away-from- +// zero: floor(x+8.5) at x=-0.5 gives 8 (rounds toward +inf), whereas +// round-half-away-from-zero would give 7. +// - SignOnly: code = sign(x0) in {-1, +1}, ignoring magnitude entirely -- +// Q1_0's actual quantizer (1-bit codes, no rounding to speak of; paired +// with ScaleAnchor::MeanAbs below, not qmax-based like every other +// anchor). +// - NearestEvenClampedHigh: round-half-to-even (not round-half-away-from- +// zero like Nearest), then clamp only the high end to qmax -- Q8_K's +// actual quantizer. GGML computes this via a magic-number float trick +// (nearest_int(), reproduced by nearest_int() below) that exploits the +// default IEEE-754 round-to-nearest-even rounding of the addition +// itself; Halide has no round-to-even primitive exposed, so the same +// bit trick is used here to match bit-for-bit. Always paired with +// ScaleAnchor::ExtremeSignedValueTwoStep below. +enum class RoundingMode { Nearest, + TruncateHalfUpWithOffset, + SignOnly, + NearestEvenClampedHigh }; + +// Same magic-number trick as GGML's static inline nearest_int() in +// src/ggml-quants.c: adding 1.5*2^23 forces the CPU's default round-to- +// nearest-even addition to round fval's fractional part, then the rounded +// integer is recovered from the float's mantissa bits. +inline Halide::Expr nearest_int(Halide::Expr fval) { + using namespace Halide; + Expr val = fval + 12582912.0f; // 1.5 * 2^23 + Expr bits = reinterpret(val); + return (bits & 0x007fffff) - 0x00400000; +} + +// How a block's scale is derived from its values -- this is a second, +// independent axis GGML varies per format, not just rounding: +// - AbsMax: scale = max(|v|) / qmax -- ordinary symmetric quantization +// (Q8_0). +// - ExtremeSignedValue: scale = -extreme / qmax, where `extreme` is the +// *signed* value with the largest magnitude in the block (ties keep the +// first-seen value, matching GGML's single left-to-right loop with a +// strict '<' comparison). This deliberately anchors the block's most +// extreme value at code -qmax, using the full negative side of an +// asymmetric signed range like [-8, 7] (Q4_0). +// - MeanAbs: scale = mean(|v|) over the block (a sum reduction divided by +// block_size, not a max reduction divided by qmax) -- Q1_0's anchor, +// always paired with RoundingMode::SignOnly. +// - ExtremeSignedValueTwoStep: mathematically the same value as +// ExtremeSignedValue (scale = -extreme/qmax), but computed as GGML's +// own two *separate* divisions -- `iscale = -qmax/extreme` first, then +// `scale = 1/iscale` -- rather than the algebraically-equivalent single +// multiply above. Floating point isn't associative, so these round +// differently in the last bit; `iscale` itself (not a fresh 1/scale +// recomputed afterward) is also what SymmetricAffineQuantize::encode() +// uses to derive codes for this anchor, to stay bit-exact with GGML's +// quantize_row_q8_K_ref -- see encode()'s `id`/`scale` computation. +// Always paired with RoundingMode::NearestEvenClampedHigh. +enum class ScaleAnchor { AbsMax, + ExtremeSignedValue, + MeanAbs, + ExtremeSignedValueTwoStep }; + +// encode(): block(kk, blk) -> {codes(kk, blk) in [-qmax, qmax-1], scale(blk)}. +// decode(): {codes, scale} -> cast(codes) * scale -- this half is +// exactly the same regardless of rounding/anchor (both Q4_0's and Q8_0's +// existing hand-written dequantize math already reduce to this one formula). +class SymmetricAffineQuantize { +public: + SymmetricAffineQuantize(int block_size, int qmax, RoundingMode rounding, ScaleAnchor anchor) + : block_size_(block_size), qmax_(qmax), rounding_(rounding), anchor_(anchor) { + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + _halide_user_assert(inputs.size() == 1) << "SymmetricAffineQuantize::encode expects one block Func\n"; + Func block = inputs[0]; + Var kk("kk"), blk("blk"); + RDom r(0, block_size_, "r"); + Func stat("symmetric_quantize_stat"), scale("symmetric_quantize_scale"), reciprocal("symmetric_quantize_reciprocal"); + auto define_extreme = [&]() { + stat(blk) = Tuple(0.0f, 0.0f); + Expr value = block(r, blk); + Expr take = abs(value) > stat(blk)[0]; + stat(blk) = Tuple(select(take, abs(value), stat(blk)[0]), + select(take, value, stat(blk)[1])); + }; + if (anchor_ == ScaleAnchor::AbsMax) { + stat(blk) = 0.0f; + stat(blk) = max(stat(blk), abs(block(r, blk))); + scale(blk) = stat(blk) / (float)qmax_; + reciprocal(blk) = select(scale(blk) != 0.0f, 1.0f / scale(blk), 0.0f); + } else if (anchor_ == ScaleAnchor::ExtremeSignedValue) { + define_extreme(); + scale(blk) = stat(blk)[1] * (-1.0f / (float)qmax_); + reciprocal(blk) = select(scale(blk) != 0.0f, 1.0f / scale(blk), 0.0f); + } else if (anchor_ == ScaleAnchor::MeanAbs) { + stat(blk) = 0.0f; + stat(blk) += abs(block(r, blk)); + scale(blk) = stat(blk) / (float)block_size_; + reciprocal(blk) = select(scale(blk) != 0.0f, 1.0f / scale(blk), 0.0f); + } else { + define_extreme(); + reciprocal(blk) = select(stat(blk)[0] == 0.0f, 0.0f, + (-1.0f * (float)qmax_) / stat(blk)[1]); + scale(blk) = select(reciprocal(blk) != 0.0f, 1.0f / reciprocal(blk), 0.0f); + } + Expr scaled = block(kk, blk) * reciprocal(blk); + Func codes("symmetric_quantize_codes"); + if (rounding_ == RoundingMode::Nearest) { + codes(kk, blk) = cast(round(scaled)); + } else if (rounding_ == RoundingMode::TruncateHalfUpWithOffset) { + // One add of the folded constant, as in GGML's `x * id + 8.5f`. + Expr raw = cast(cast(scaled + ((float)qmax_ + 0.5f))); + codes(kk, blk) = cast(min(raw, 2 * qmax_ - 1) - qmax_); + } else if (rounding_ == RoundingMode::SignOnly) { + codes(kk, blk) = cast(select(block(kk, blk) >= 0.0f, 1, -1)); + } else { + codes(kk, blk) = cast(min(qmax_, nearest_int(scaled))); + } + return {codes, scale}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + _halide_user_assert(encoded.size() == 2) << "SymmetricAffineQuantize::decode expects codes and scale\n"; + Var kk("kk"), blk("blk"); + Func dequantized("symmetric_dequantized"); + dequantized(kk, blk) = cast(encoded[0](kk, blk)) * encoded[1](blk); + return {dequantized}; + } + + /** (within, block) floats <-> int8 codes and one float scale per block. + * + * The codes' range is a guarantee for finite inputs: [-qmax, qmax] for + * round-to-nearest, [-qmax, qmax - 1] for TruncateHalfUpWithOffset, and + * [-1, 1] for SignOnly. It is not declared for MeanAbs scales (except + * SignOnly), where a code may be as large as the block size. */ + Halide::ApproximationSignature signature() const { + std::optional codes; + if (rounding_ == RoundingMode::SignOnly) { + codes = Halide::ApproximationRange(-1, 1); + } else if (anchor_ != ScaleAnchor::MeanAbs) { + const int hi = rounding_ == RoundingMode::TruncateHalfUpWithOffset ? qmax_ - 1 : qmax_; + codes = Halide::ApproximationRange(std::max(-qmax_, -128), std::min(hi, 127)); + } + return {{{"block", Halide::Float(32), 2}}, + {{"codes", Halide::Int(8), 2, codes}, {"scale", Halide::Float(32), 1}}}; + } + + /** For finite inputs, a bound on |decode(encode(x)) - x| in units of the + * block's |scale|, when the scale is set by the block's extreme value + * (AbsMax, ExtremeSignedValue, ExtremeSignedValueTwoStep) and the codes + * are not sign-only: + * + * - Nearest and NearestEvenClampedHigh round to the nearest code, so the + * error is at most half a step: |scale| / 2; + * - TruncateHalfUpWithOffset also clamps the top code (a value that + * scales to +qmax lands on qmax - 1), so the bound is a full step: + * |scale|. + * + * Both include a relative slack of qmax * 2^-21 covering the float + * rounding of the reciprocal, the scaled value, and the dequantizing + * product. No bound is declared for MeanAbs scales or SignOnly. */ + Halide::Func error_bound(const std::vector &inputs, const std::vector &encoded) const { + if (anchor_ == ScaleAnchor::MeanAbs || rounding_ == RoundingMode::SignOnly) { + return Halide::Func(); + } + using namespace Halide; + _halide_user_assert(inputs.size() == 1 && encoded.size() == 2) + << "SymmetricAffineQuantize::error_bound expects one input and two encoded Funcs\n"; + const double steps = rounding_ == RoundingMode::TruncateHalfUpWithOffset ? 1.0 : 0.5; + const double factor = steps + qmax_ * (1.0 / (1 << 21)); + Var kk("kk"), blk("blk"); + Func bound("symmetric_quantize_error_bound"); + bound(kk, blk) = abs(cast(encoded[1](blk))) * Expr(factor); + return bound; + } + +private: + int block_size_, qmax_; + RoundingMode rounding_; + ScaleAnchor anchor_; +}; + +// How AffineQuantize rounds+truncates code = round((x-min)*id) into its +// final representable range -- a different formula than SymmetricAffineQuantize's +// RoundingMode, and not a variation on it: there's no centering/offset here +// (codes are naturally unsigned starting at 0), and GGML's two affine legacy +// formats don't even agree on whether to clamp at all: +// - ClampedInt8: (int8_t)(v+0.5f), then an explicit min(.., levels) -- +// Q4_1's exact formula. +// - UnclampedUint8: (uint8_t)(v+0.5f) directly, no further clamp -- Q5_1's +// exact formula. GGML's own reference genuinely doesn't clamp this one +// (verified against quantize_row_q5_1_ref); reproduced faithfully since +// quantize output is checked bit-exact. +enum class AffineRounding { ClampedInt8, + UnclampedUint8 }; + +// encode(): block(kk, blk) -> {codes(kk, blk) in [0, levels], scale(blk), +// min(blk)}. decode(): {codes, scale, min} -> cast(codes)*scale + min. +// The min-max (not max-abs) scale derivation is what makes this "affine" +// rather than "symmetric" -- every value in a block is representable, not +// just those centered on zero, at the cost of needing a second per-block +// float (Q4_1/Q5_1's 'm'). +class AffineQuantize { +public: + AffineQuantize(int block_size, int levels, AffineRounding rounding) + : block_size_(block_size), levels_(levels), rounding_(rounding) { + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + Func block = inputs[0]; // block(kk, blk) + Var kk("kk"), blk("blk"); + RDom r(0, block_size_, "r"); + + // Plain min/max reduction -- unlike SymmetricAffineQuantize's + // ScaleAnchor::ExtremeSignedValue, min and max are independent here, + // forming an affine (not centered-on-zero) range. + Func stat("affine_quantize_minmax"); + stat(blk) = Tuple(std::numeric_limits::max(), std::numeric_limits::lowest()); + Expr v = block(r, blk); + stat(blk) = Tuple(min(stat(blk)[0], v), max(stat(blk)[1], v)); + + Func scale("affine_quantize_scale"); + scale(blk) = (stat(blk)[1] - stat(blk)[0]) / (float)levels_; + Func minv("affine_quantize_min"); + minv(blk) = stat(blk)[0]; + + Expr id = select(scale(blk) != 0.0f, 1.0f / scale(blk), 0.0f); + Expr x0 = (block(kk, blk) - minv(blk)) * id; + + Func codes("affine_quantize_codes"); + if (rounding_ == AffineRounding::ClampedInt8) { + Expr raw = cast(cast(x0 + 0.5f)); + codes(kk, blk) = cast(min(raw, levels_)); + } else { + codes(kk, blk) = cast(cast(x0 + 0.5f)); + } + + return {codes, scale, minv}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + Func codes = encoded[0], scale = encoded[1], minv = encoded[2]; + Var kk("kk"), blk("blk"); + Func dequantized("affine_dequantized_am"); + dequantized(kk, blk) = cast(codes(kk, blk)) * scale(blk) + minv(blk); + return {dequantized}; + } + +private: + int block_size_, levels_; + AffineRounding rounding_; +}; + +// --------------------------------------------------------------------------- +// 2b. Units for Q5_0's split 5-bit code: an additive low/high radix split and a +// one-bit-per-element alphabet pack. +// --------------------------------------------------------------------------- + +namespace detail { + +inline std::vector component_vars(int dimensions, const std::string &prefix) { + std::vector vars; + vars.reserve(dimensions); + for (int i = 0; i < dimensions; ++i) { + vars.emplace_back(prefix + std::to_string(i)); + } + return vars; +} + +inline std::vector component_exprs(const std::vector &vars) { + return std::vector(vars.begin(), vars.end()); +} + +// The dimensionality of the single port in `inputs`, if there is exactly one. +inline std::optional single_port_dimensions(const Halide::ApproximationPorts &inputs) { + if (inputs.size() == 1) { + return inputs[0].dimensions; + } + return std::nullopt; +} + +} // namespace detail + +/** Pack a fixed vector containing exactly two values into an integer word. + * Decode uses an embedded 8x256 byte-expansion LUT, allowing each source byte + * to expand through one contiguous eight-byte load. */ +template +class BinaryAlphabetPack { +public: + BinaryAlphabetPack(int vector_size, Halide::Type word_type, Value zero_value, Value one_value) + : vector_size_(vector_size), word_type_(word_type), zero_value_(zero_value), one_value_(one_value), + expansion_(8, 256) { + _halide_user_assert(word_type_.is_uint() && word_type_.bits() >= vector_size_) + << "BinaryAlphabetPack word is too small for its vector\n"; + _halide_user_assert(vector_size_ > 0 && vector_size_ % 8 == 0) + << "BinaryAlphabetPack vector size must be a positive multiple of eight\n"; + for (int byte = 0; byte < 256; ++byte) { + for (int bit = 0; bit < 8; ++bit) { + expansion_(bit, byte) = (byte & (1 << bit)) ? one_value_ : zero_value_; + } + } + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + _halide_user_assert(inputs.size() == 1 && inputs[0].dimensions() >= 2) + << "BinaryAlphabetPack::encode requires (element, record...)\n"; + Func values = inputs[0]; + int records_n = values.dimensions() - 1; + std::vector records = detail::component_vars(records_n, "record"); + std::vector record_args = detail::component_exprs(records); + RDom bit(0, vector_size_, "binary_bit"); + std::vector value_args = record_args; + value_args.insert(value_args.begin(), bit); + Expr value = values(value_args); + Func word("binary_alphabet_word"); + word(records) = cast(word_type_, 0); + word(records) = word(record_args) | + select(value == cast(one_value_), + cast(word_type_, 1) << bit, cast(word_type_, 0)); + return {word}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + _halide_user_assert(encoded.size() == 1 && encoded[0].types() == std::vector{word_type_}) + << "BinaryAlphabetPack::decode word type mismatch\n"; + Func word = encoded[0]; + std::vector records = detail::component_vars(word.dimensions(), "record"); + std::vector record_args = detail::component_exprs(records); + Var element("element"); + Expr bits = word(record_args); + Expr byte = cast((bits >> ((element / 8) * 8)) & 0xff); + std::vector args = records; + args.insert(args.begin(), element); + Func values("binary_alphabet_values"); + values(args) = expansion_(element % 8, byte); + return {values}; + } + + /** (element, record...) values <-> one word per record. */ + Halide::ApproximationSignature signature(const Halide::ApproximationPorts &inputs) const { + if (inputs.size() > 1) { + return Halide::ApproximationSignature::unknown(inputs); + } + std::optional dims = detail::single_port_dimensions(inputs); + return {{{inputs.empty() ? "values" : inputs[0].name, Halide::type_of(), dims}}, + {{"word", word_type_, dims ? std::optional(*dims - 1) : std::nullopt}}}; + } + +private: + int vector_size_; + Halide::Type word_type_; + Value zero_value_, one_value_; + Halide::Buffer expansion_; +}; +/** Split a signed code into an unsigned low digit and an additive weighted + * high contribution. The second parameter recenters the signed code before + * taking the low digit; decode is simply `code = low + high`. */ +class AdditiveRadixSplit { +public: + AdditiveRadixSplit(int radix, int offset) + : radix_(radix), offset_(offset) { + _halide_user_assert(radix_ > 1 && offset_ >= 0) << "Invalid AdditiveRadixSplit parameters\n"; + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + _halide_user_assert(inputs.size() == 1) << "AdditiveRadixSplit::encode expects one code Func\n"; + Func code = inputs[0]; + std::vector args = detail::component_vars(code.dimensions(), "code"); + std::vector call_args = detail::component_exprs(args); + Expr value = cast(code(call_args)); + Expr low_value = (value + offset_) % radix_; + Func low("additive_radix_low"), high("additive_radix_high"); + low(args) = cast(low_value); + high(args) = cast(value - low_value); + return {low, high}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + _halide_user_assert(encoded.size() == 2 && encoded[0].dimensions() == encoded[1].dimensions()) + << "AdditiveRadixSplit::decode expects low and high contributions\n"; + std::vector args = detail::component_vars(encoded[0].dimensions(), "code"); + std::vector call_args = detail::component_exprs(args); + Func code("additive_radix_code"); + code(args) = cast(cast(encoded[0](call_args)) + + cast(encoded[1](call_args))); + return {code}; + } + + /** A code <-> its unsigned low digit and signed high contribution. + * + * The input range is the precondition for exactness, found by trying every + * int8 code (decode returns an int8): the codes for which the low digit + * fits a byte and the high contribution fits an int8. The output ranges + * are the digits those codes produce. No ranges are declared if the valid + * codes are empty or not contiguous. */ + Halide::ApproximationSignature signature(const Halide::ApproximationPorts &inputs) const { + if (inputs.size() > 1) { + return Halide::ApproximationSignature::unknown(inputs); + } + std::optional dims = detail::single_port_dimensions(inputs); + Analysis a = analyze(); + return {{{inputs.empty() ? "code" : inputs[0].name, std::nullopt, dims, a.valid ? std::optional(a.code) : std::nullopt}}, + {{"low", Halide::UInt(8), dims, a.valid ? std::optional(a.low) : std::nullopt}, + {"high", Halide::Int(8), dims, a.valid ? std::optional(a.high) : std::nullopt}}}; + } + + /** Exact for codes in the declared input range. */ + bool lossless() const { + return analyze().valid; + } + +private: + int radix_, offset_; + + struct Analysis { + bool valid = false; + Halide::ApproximationRange code, low, high; + }; + + Analysis analyze() const { + Analysis a; + int count = 0, first = 0, last = 0; + for (int v = -128; v <= 127; ++v) { + int low = ((v + offset_) % radix_ + radix_) % radix_; + int high = v - low; + if (low > 255 || high < -128 || high > 127) { + continue; + } + if (count++ == 0) { + first = v; + a.low = Halide::ApproximationRange(low, low); + a.high = Halide::ApproximationRange(high, high); + } + last = v; + a.low = Halide::ApproximationRange(std::min(a.low.lo, low), std::max(a.low.hi, low)); + a.high = Halide::ApproximationRange(std::min(a.high.lo, high), std::max(a.high.hi, high)); + } + a.valid = count > 0 && count == last - first + 1; + a.code = Halide::ApproximationRange(first, last); + return a; + } +}; +// --------------------------------------------------------------------------- +// Elementwise storage policies, as Halide::Pointwise units. +// --------------------------------------------------------------------------- + +// A float32 value stored as a float16 (decode widens it back). strict_float +// keeps the storage rounding from being folded away when the encode side is +// inlined into its consumers. Lossy, so no losslessness is claimed. +inline Halide::Pointwise fp16_storage() { + using namespace Halide; + return Pointwise{"storage_cast_stored", "storage_cast_decoded", + [](Expr x) { return strict_float(cast(x)); }, + [](Expr x) { return strict_float(cast(x)); }, + "cast"} + .with_types(Float(32), Float(16)); +} + +// Translate signed int8 codes in [-offset, 15 - offset] into the unsigned +// nibbles [0, 15] that PlanarFieldPack{4, ...} stores, e.g. q4_0's [-8, 7] +// for offset 8. This is representation policy, not bit packing. +inline Halide::Pointwise nibble_offset(int offset) { + using namespace Halide; + Expr shift = Internal::make_const(Int(64), (int64_t)offset); + return Pointwise{"additive_offset_stored", "additive_offset_decoded", + [shift](Expr x) { return cast(cast(x) + shift); }, + [shift](Expr x) { return cast(cast(x) - shift); }, + "offset"} + .with_types(Int(8), UInt(8)) + .with_ranges(ApproximationRange(-offset, 15 - offset), ApproximationRange(0, 15)) + .with_lossless(); +} + +// --------------------------------------------------------------------------- +// 3a. Per-field bit packing. +// --------------------------------------------------------------------------- + +// Assemble a little-endian integer word starting at bytes(base, blk). The +// uint32 specialization below is the on-disk byte order shared by F32Pack's +// scale and IQ3_XXS's/IQ2_XXS's aux32. +inline Halide::Expr le_uint(Halide::Func bytes, Halide::Expr base, Halide::Var blk, int byte_count) { + using namespace Halide; + Expr result = cast(0); + for (int i = 0; i < byte_count; i++) { + result = result | (cast(bytes(base + i, blk)) << (8 * i)); + } + return result; +} + +inline Halide::Expr le_u32(Halide::Func bytes, Halide::Expr base, Halide::Var blk) { + using namespace Halide; + return le_uint(bytes, base, blk, 4); +} + +// le_uint's encode-side mirror: byte `byte_idx` (little-endian) of `bits`. +// Widening to uint32 first makes one expression work for every word size +// used below (Fp16Pack's 16-bit word, F32Pack's 32-bit one, Int16Pack's +// per-group 16-bit one) -- shifting a narrower type by up to 24 bits would +// silently truncate instead. +inline Halide::Expr word_to_le_byte(Halide::Expr bits, Halide::Expr byte_idx) { + using namespace Halide; + return cast((cast(bits) >> (cast(byte_idx) * 8)) & 0xff); +} + +// le_uint's encode-side-agnostic generalization: a little-endian 16-bit word +// from two bytes, but via an arbitrary per-byte accessor (`byte_at(0)` = low +// byte, `byte_at(1)` = high byte) instead of a fixed `Func bytes` indexed at +// (base+i, blk) -- so it also covers a Halide::_-general decoder's +// bytes(offset, blk, _) reads, or a byte pair that isn't a plain Func index +// at all (e.g. IQ1_M's per-scale-word accumulation further below). Replaces +// the hand-rolled "lo | (hi << 8)" pattern that used to be written out at +// each call site. +template +inline Halide::Expr le_u16(ByteAt byte_at) { + using namespace Halide; + Expr lo = cast(byte_at(0)); + Expr hi = cast(byte_at(1)); + return cast(lo | (hi << 8)); +} + +// encode(float scale) -> 2 bytes; decode(2 bytes) -> float. Matches every +// format's existing fp16 delta-byte code (byte0 = bits&0xff, byte1 = +// (bits>>8)&0xff). +class Fp16Pack { +public: + Halide::Func encode(const Halide::Func &scale) const { + using namespace Halide; + // scale(blk) + Var byte("byte"), blk("blk"); + Expr bits = reinterpret(cast(scale(blk))); + Func bytes("fp16_pack_bytes"); + bytes(byte, blk) = word_to_le_byte(bits, byte); + return bytes; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + // bytes(byte, blk[, ...]), byte in [0, 2) + Var blk("blk"); + Expr bits = le_u16([&](int i) { return bytes(i, blk, _); }); + Func scale("fp16_pack_scale"); + scale(blk, _) = cast(reinterpret(bits)); + return scale; + } +}; + +// encode(float scale) -> 4 bytes (plain IEEE-754 binary32, little-endian); +// decode(4 bytes) -> float -- Q8_K's scale format, the one format here whose +// delta is a full float, not fp16 (matching block_q8_K's `float d;`, not +// GGML's usual `ggml_fp16_t d;`). +class F32Pack { +public: + Halide::Func encode(const Halide::Func &scale) const { + using namespace Halide; + // scale(blk) + Var byte("byte"), blk("blk"); + Expr bits = reinterpret(scale(blk)); + Func bytes("f32_pack_bytes"); + bytes(byte, blk) = word_to_le_byte(bits, byte); + return bytes; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + // bytes(byte, blk[, ...]), byte in [0, 4) + Var blk("blk"); + // Dimension-general via Halide::_ (matches Fp16Pack): any trailing + // "lane" dims -- e.g. a repack weight's columns -- ride through, so this + // pack can decode a repack scale field, not just a plain (byte, blk) one. + Expr b0 = cast(bytes(0, blk, _)), b1 = cast(bytes(1, blk, _)); + Expr b2 = cast(bytes(2, blk, _)), b3 = cast(bytes(3, blk, _)); + Func scale("f32_pack_scale"); + scale(blk, _) = reinterpret(b0 | (b1 << 8) | (b2 << 16) | (b3 << 24)); + return scale; + } +}; + +// encode(int16 values(g, blk)) -> 2 little-endian bytes each, byte `2g`/ +// `2g+1` -- e.g. Q8_K's bsums[16] int16 array (one value per 16-element +// group, not one per block). +class Int16Pack { +public: + Halide::Func encode(const Halide::Func &values) const { + using namespace Halide; + // values(g, blk) + Var byte_idx("byte_idx"), blk("blk"); + Expr g = byte_idx / 2; + Expr bits = reinterpret(values(g, blk)); + Func bytes("int16_pack_bytes"); + bytes(byte_idx, blk) = word_to_le_byte(bits, byte_idx % 2); + return bytes; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + // bytes(byte_idx, blk), byte_idx in [0, 2*num_groups) + Var g("g"), blk("blk"); + Expr bits = le_u16([&](int i) { return bytes(2 * g + i, blk); }); + Func values("int16_pack_values"); + values(g, blk) = reinterpret(bits); + return values; + } +}; + +// decode(1 byte) -> float, an E8M0 power-of-two exponent (GGML's MXFP4/NVFP4 +// scale format). Reproduces ggml_e8m0_to_fp32_half's exact bit construction: +// d = 2^(e-128) for every e in [0, 255], computed via a subnormal-exploiting +// shift trick for e<2 instead of a normal exponent-field write (both branches +// compute the same uniform 2^(e-128); see the comment inline). Decode-only: +// MXFP4 quantize is extern-delegated (see ExternQuantize). +class E8M0Pack { +public: + std::vector encode(const std::vector &) const { + _halide_user_error << "E8M0Pack is decode-only -- quantize is deferred to an ExternQuantize.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &byte) const { + using namespace Halide; + // byte(byte_idx, blk[, ...]), byte_idx in [0, 1) + Var blk("blk"); + // Dimension-general via Halide::_ (matches Fp16Pack/F32Pack) so this pack + // can decode a repack weight's E8M0 scale header, columns riding through. + Expr e = cast(byte(0, blk, _)); + Expr bits = select(e < 2, cast(0x00200000) << e, (e - 1) << 23); + Func scale("e8m0_pack_scale"); + scale(blk, _) = reinterpret(bits); + return scale; + } +}; + +// decode(1 byte per sub-block) -> float(sub, blk), a UE4M3 unsigned +// 4-exponent/3-mantissa float (GGML's NVFP4 per-sub-block scale format). +// Reproduces ggml_ue4m3_to_fp32_half's exact construction: subnormal +// (exp==0) is man/512, normal is (1+man/8)*2^(exp-7), both halved; byte 0x00 +// or 0x7f (GGML's NVFP4 zero/sentinel bytes) decode to 0. Unlike +// Fp16Pack/E8M0Pack (exactly one scale value per block), this is meant to be +// used with LinearDequant's per-sub-block (sub_size > 0) mode: the Funcs here are indexed by +// `sub` directly (one byte each). Decode-only: NVFP4 quantize is +// extern-delegated (see ExternQuantize). +class UE4M3Pack { +public: + std::vector encode(const std::vector &) const { + _halide_user_error << "UE4M3Pack is decode-only -- quantize is deferred to an ExternQuantize.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &byte) const { + using namespace Halide; + // byte(sub, blk) + Var sub("sub"), blk("blk"); + Expr ue = byte(sub, blk); // codespell:ignore ue + Expr is_zero = (ue == 0) || (ue == 0x7f); // codespell:ignore ue + Expr exp_ = cast(cast(ue) >> 3) & 0xf; // codespell:ignore ue + Expr man_ = cast(ue) & 0x7; // codespell:ignore ue + Expr raw = select(exp_ == 0, + cast(man_) / 512.0f, + (1.0f + cast(man_) / 8.0f) * pow(2.0f, cast(exp_ - 7))); + Func scale("ue4m3_pack_scale"); + scale(sub, blk) = select(is_zero, 0.0f, raw * 0.5f); + return scale; + } +}; + +// The four scalar-scale on-disk formats above (Fp16Pack/F32Pack/E8M0Pack/ +// UE4M3Pack), named so call sites that need to pick one at runtime (the +// repack weight schemes' RepackWeightScale used to be a bespoke copy of +// exactly this same enum; make_codebook_scheme's scale_pack/scale_bytes +// parameter pair was the other) can hand a single value across instead of a +// (Approximation, int width) pair kept in sync by hand. +enum class ScaleFormat { Fp16, + F32, + E8M0, + UE4M3 }; + +inline Halide::Approximation make_scale_pack(ScaleFormat fmt) { + switch (fmt) { + case ScaleFormat::Fp16: + return Fp16Pack(); + case ScaleFormat::F32: + return F32Pack(); + case ScaleFormat::E8M0: + return E8M0Pack(); + case ScaleFormat::UE4M3: + return UE4M3Pack(); + } + _halide_internal_error << "unreachable ScaleFormat\n"; +} + +// On-disk byte width of `num_scales` consecutive values in `fmt` (e.g. +// NVFP4's 4 per-sub-block UE4M3 bytes, or a repack weight's n_cols-wide +// per-column scale header). +inline int scale_width(ScaleFormat fmt, int num_scales = 1) { + switch (fmt) { + case ScaleFormat::Fp16: + return 2 * num_scales; + case ScaleFormat::F32: + return 4 * num_scales; + case ScaleFormat::E8M0: + return num_scales; + case ScaleFormat::UE4M3: + return num_scales; + } + _halide_internal_error << "unreachable ScaleFormat\n"; +} + +// The shared byte<->field decomposition behind every uniform-width, +// non-base-3 bit-packed field in this file (nibble packs, 2-bit packs, and +// the K-quants' rotating out-of-band high-bit arrays): a flat per-block +// index kk factors as +// kk = outer*(plane_count*pos_count) + plane*pos_count + pos +// and each byte packs `plane_count` fixed-width `field_bits`-bit fields +// ("planes") -- byte `outer*pos_count + pos` holds element (outer, plane, pos) +// at bit-shift `plane*field_bits`. This one decomposition is exactly (verified +// by hand, not just formally) GGML's addressing for all of the following, each +// just a different (field_bits, plane_count, pos_count) instantiated directly +// at its call site below rather than through a named subclass, since the +// class itself is the entire "component" here -- a named wrapper would only +// rename a handful of integers, not remove any duplication: +// - Nibble-packed formats (Q4_0, Q4_1, Q5_0's low bits, Q4_K, NVFP4, ...): +// 4-bit fields, 2 planes (low/high nibble), pos_count = window_size/2, +// outer = which independent sub-block "window" (1 for a plain whole-block +// low-half/high-half split, e.g. Q4_0/IQ4_NL; >1 for NVFP4/Q4_K/Q5_K's +// independent windows) -- GGML's actual convention for every nibble- +// packed format, not just the more obvious "adjacent pair per byte" +// layout. +// - TQ2_0/Q2_K/Q3_K/Q6_K's 2-bit fields: 2-bit fields, 4 planes, pos_count +// = block_size/8 (= window_bytes), outer = which half of the block -- +// verified by hand against each format's original dequantize math (e.g. +// TQ2_0's `half = gi/half_block; l = local/window_bytes; m = +// local%window_bytes; byte_idx = half*window_bytes + m; shift = l*2`); +// NOT a plain "4 consecutive elements per byte" packing. +// - Q3_K's hmask / Q5_K's qh out-of-band high-bit arrays: 1-bit fields, +// num_windows planes, pos_count = window_size, outer always 0 (the whole +// array is one window_size*num_windows group) -- the "rotating bit +// position" scheme GGML uses for an extra high bit per element (their own +// source expresses this via a `half`/`iter`-based case split instead, but +// it's the same addressing). Always paired with a lower-bit code via +// CombineBits, never used standalone. +// `qmax` recenters already-decoded codes before splitting (0 leaves the raw +// field, e.g. a lookup-table index or a single out-of-band bit -- always 0 +// for the two cases above). encode() OR-accumulates the planes into each +// byte via an RDom, exactly like BitPack's qh_accum-style OR-reduction. +// +// `plane_axis` (default false) selects an alternate decode that keeps `plane` +// as an explicit LEADING output axis instead of folding it into a flat kk: +// fields(plane, pos, blk, _) = (bytes(pos, blk, _) >> plane*field_bits) & mask +// This is the shape a *combined* (scale, min) field wants -- plane 0 = scale, +// plane 1 = min -- so LinearDequant can read both from one func by its +// plane index. It's dimension-general (Halide::_) and needs neither `outer` nor +// `qmax` (raw unsigned fields), so it is exactly Q2_K's per-sub-block nibble- +// pair (scale, min) byte array, with no bespoke leaf. Decode-only in this mode +// (only ever an extern-delegated TrustedInverse decoder stage). +class PlanarBitPack { +public: + // plane_count is always 8/field_bits in every instantiation here (a byte + // packs exactly 8 bits' worth of same-width fields, full stop), so it's + // derived rather than taken as its own parameter -- see nibble_pack/ + // crumb_pack/rotating_bit_pack/le_bit_pack below for the named-shape + // constructors most call sites should use instead of this directly. + PlanarBitPack(int field_bits, int pos_count, int qmax = 0, bool plane_axis = false, + int value_scale = 1) + : field_bits_(field_bits), plane_count_(8 / field_bits), pos_count_(pos_count), qmax_(qmax), + plane_axis_(plane_axis), value_scale_(value_scale) { + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + if (plane_axis_) { + _halide_user_error << "PlanarBitPack plane-axis mode is decode-only " + "(only an extern-delegated TrustedInverse decoder stage).\n"; + return {}; + } + Func codes = inputs[0]; + Var byte_idx("byte_idx"), blk("blk"); + int group = plane_count_ * pos_count_; + Expr outer = byte_idx / pos_count_; + Expr pos = byte_idx % pos_count_; + + RDom rp(0, plane_count_, "rp"); + Expr kk = outer * group + rp * pos_count_ + pos; + Expr field = cast((cast(codes(kk, blk)) + qmax_) / value_scale_) & + ((1u << field_bits_) - 1); + Func bytes("planar_bit_pack_bytes"); + bytes(byte_idx, blk) = cast(0); + bytes(byte_idx, blk) = bytes(byte_idx, blk) | cast(field << (rp * field_bits_)); + return {bytes}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + Var kk("kk"), blk("blk"); + if (plane_axis_) { + Var plane("plane"), pos("pos"); + Expr field = (cast(bytes(pos, blk, _)) >> (cast(plane) * field_bits_)) & + ((1u << field_bits_) - 1); + Func fields("planar_bit_pack_fields"); + fields(plane, pos, blk, _) = cast(field); + return fields; + } + // Lane-general via Halide::_ (see the DIMENSION / WILDCARD CONVENTION + // in the file's top comment), like the plane-axis mode above. + int group = plane_count_ * pos_count_; + Expr outer = kk / group; + Expr rem = kk % group; + Expr plane = rem / pos_count_; + Expr pos = rem % pos_count_; + Expr byte_idx = outer * pos_count_ + pos; + Expr byte = cast(bytes(byte_idx, blk, _)); + Expr field; + if (field_bits_ == 1) { + // A single-bit field (Q5_0/Q5_1's per-element high bit). Read it from + // a compile-time byte->bits expansion table (ggml's table_b2b idea): + // b2b(bit, byte) = (byte >> bit) & 1. The alternative arithmetic form + // `(byte >> plane) & 1` lowers to a per-lane byte-broadcast transpose + // plus mask (~24 NEON ops per 16 lanes), which is the whole + // q5_x-vs-q4_x gap. A Buffer<> is embedded in the binary as constant + // data, so this is a pure lookup, no runtime input. + static const Halide::Buffer b2b = [] { + Halide::Buffer t(8, 256); + t.set_min(0, 0); + for (int by = 0; by < 256; by++) { + for (int bit = 0; bit < 8; bit++) { + t(bit, by) = (by >> bit) & 1; + } + } + return t; + }(); + static const Halide::Buffer b2b_shifted = [] { + Halide::Buffer t(8, 256); + t.set_min(0, 0); + for (int by = 0; by < 256; by++) { + for (int bit = 0; bit < 8; bit++) { + t(bit, by) = ((by >> bit) & 1) << 4; + } + } + return t; + }(); + static const Halide::Buffer b2b_shifted_signed = [] { + Halide::Buffer t(8, 256); + t.set_min(0, 0); + for (int by = 0; by < 256; by++) { + for (int bit = 0; bit < 8; bit++) { + t(bit, by) = (((by >> bit) & 1) << 4) - 16; + } + } + return t; + }(); + const Expr bit_index = cast(plane * field_bits_); + const Expr byte_value = cast(byte); + if (value_scale_ == 16 && qmax_ == 16) { + field = b2b_shifted_signed(bit_index, byte_value); + } else if (value_scale_ == 16 && qmax_ == 0) { + field = b2b_shifted(bit_index, byte_value); + } else { + field = cast(b2b(bit_index, byte_value)) * value_scale_ - qmax_; + } + } else { + field = cast((byte >> (plane * field_bits_)) & ((1u << field_bits_) - 1)) * value_scale_ - qmax_; + } + Func codes("planar_bit_pack_codes"); + codes(kk, blk, _) = cast(field); + return codes; + } + +private: + int field_bits_, plane_count_, pos_count_, qmax_; + bool plane_axis_; + int value_scale_; +}; + +// Named PlanarBitPack shapes for the four bit-widths this file actually +// instantiates, so a call site reads as "a nibble pack over this many +// elements" instead of a bare (field_bits, pos_count) pair the reader has to +// re-derive the meaning of. `window` is the element span PlanarBitPack's own +// class comment calls "window_size": for nibble_pack/crumb_pack, the size of +// one independently low/high- (or low/high-2-bit-) split group (pos_count is +// window/2 or window/4, one plane's share of it); for rotating_bit_pack, the +// single rotating-bit-position span itself (pos_count = window directly, +// since GGML's hmask/qh arrays have exactly one such group covering the +// *whole* field, not several independent windows). +inline Halide::Approximation nibble_pack(int window, int qmax = 0) { + return PlanarBitPack(4, window / 2, qmax); +} +inline Halide::Approximation crumb_pack(int window, int qmax = 0) { + return PlanarBitPack(2, window / 4, qmax); +} +inline Halide::Approximation rotating_bit_pack(int window, int qmax = 0) { + return PlanarBitPack(1, window, qmax); +} +// The Stage-2 qh addressing (make_code_pack's code_bits==5 case): one flat +// bit per element, byte kk/8 at shift kk%8 -- PlanarBitPack{1, 1} regardless +// of block size (pos_count is always 1; there's no "window" to parameterize). +inline Halide::Approximation le_bit_pack(int value_scale = 1, int qmax = 0) { + return PlanarBitPack(1, 1, qmax, false, value_scale); +} + +// GGML's Q5_0/Q5_1 5-bit code split (a 4-bit low nibble plus a 5th high bit, +// OR-accumulated one bit per element into a 32-bit little-endian word) used +// to be a bespoke FiveBitPack class here. It's deleted: verified by hand that +// it is exactly the K-quant "combined bit code" shape (CombineBits, section +// 5 below) already used for Q3_K/Q5_K/Q6_K's own adjacent {high-bit array; +// low-bits array} regions -- qh's bit `kk` is element kk's own high bit, +// i.e. le_bit_pack()'s addressing (byte kk/8, shift kk%8; GGML's own code +// computes this via a low/high-half split for scalar-loop efficiency -- e.g. +// `(qh >> (byte_idx+12)) & 0x10` for the high half -- but that's an +// equivalent, more roundabout way of writing the same fact used directly +// here), and the nibble half is exactly nibble_pack(block_size) (byte b +// holds elements b and b+block_size/2 at shifts 0/4). See make_code_pack's +// code_bits==5 case below, which assembles this +// via make_combined_bit_codec exactly as make_q5_k_scheme does for its own +// adjacent {qh; qs} region. + +// encode(codes(kk, blk) signed in {-1, +1}) -> block_size/8 bytes; decode +// reverses it. Packs one sign bit per element, 8 elements per byte, bit +// `kk % 8` of byte `kk / 8` set when code is +1 -- Q1_0's layout (paired +// with RoundingMode::SignOnly/ScaleAnchor::MeanAbs above). Accumulates via +// an OR-reduction exactly like PlanarBitPack::encode's per-byte accumulation, +// just one full byte's worth of bits at a time instead of a per-plane field. +class BitPack { +public: + Halide::Func encode(const Halide::Func &codes) const { + using namespace Halide; + Var byte_idx("byte_idx"), blk("blk"); + + Func bit("bit_pack_bit"); + Var kk("kk"); + bit(kk, blk) = cast(select(codes(kk, blk) > 0, 1, 0)); + + RDom rb(0, 8, "rb"); + Func bytes("bit_pack_bytes"); + bytes(byte_idx, blk) = cast(0); + bytes(byte_idx, blk) = bytes(byte_idx, blk) | cast(bit(byte_idx * 8 + rb, blk) << rb); + + return bytes; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + // bytes(byte_idx, blk), byte_idx in [0, block_size/8) + Var kk("kk"), blk("blk"); + Expr byte_idx = kk / 8; + Expr bit_off = kk % 8; + Expr bit = (cast(bytes(byte_idx, blk)) >> bit_off) & 1u; + Func codes("bit_pack_codes"); + codes(kk, blk) = cast(select(bit != 0, 1, -1)); + return codes; + } +}; + +// encode(codes(kk, blk) signed int8) -> 1 code per byte, same shape (the +// identity-shaped case PlanarBitPack's nibble/2-bit packing doesn't cover, +// used by e.g. Q8_0). Unlike PlanarBitPack, this formula has no precondition +// on kk's range (reinterpret is valid for any kk) -- it doesn't know or care +// what block_size is; bounds propagate backward from whatever actually +// consumes it. +class BytePack { +public: + Halide::Func encode(const Halide::Func &codes) const { + using namespace Halide; + Var kk("kk"), blk("blk"); + Func bytes("byte_pack_bytes"); + bytes(kk, blk) = reinterpret(codes(kk, blk)); + return bytes; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + Var kk("kk"), blk("blk"); + // Lane-general via Halide::_ (see the DIMENSION / WILDCARD CONVENTION + // in the file's top comment): trailing lane dims ride through. + Func codes("byte_pack_codes"); + codes(kk, blk, _) = reinterpret(bytes(kk, blk, _)); + return codes; + } +}; + +// decode(52 bytes {qs[48]; qh[4]}) -> codes(kk, blk), the raw base-3 digit in +// [0, 3) -- GGML's TQ1_0 packing: 256 elements in 3 sections (a 160-element +// and an 80-element run at 5 trits/byte, then a 16-element run at 4 real +// trits/byte, its 5th digit always 0). decode() reverses the ceiling-division +// packing (`byte = ceil(digit_number * 256 / 243)`) via the same +// multiply-truncate-rescale trick tq1_0_generators.cpp hand-rolled: extracting +// digit `n` needs multiplier `3^n`, `n` from the most-significant digit. Codes +// feed a codebook index directly (TQ2_0's {-1, 0, 1, unused} convention). +// Decode-only: TQ1_0 quantize is extern-delegated. +class TritPack { +public: + std::vector encode(const std::vector &) const { + _halide_user_error << "TritPack is decode-only -- quantize is deferred to an ExternQuantize.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + // bytes(byte_idx, blk), byte_idx in [0, 52) + Var kk("kk"), blk("blk"); + + Expr n_a = kk / 32; + Expr byte_a = kk % 32; + + Expr local_b = kk - 160; + Expr n_b = local_b / 16; + Expr byte_b = 32 + local_b % 16; + + Expr local_c = kk - 240; + Expr n_c = local_c / 4; + Expr byte_c = 48 + local_c % 4; + + Expr n = select(kk < 160, n_a, select(kk < 240, n_b, n_c)); + Expr byte_abs = select(kk < 160, byte_a, select(kk < 240, byte_b, byte_c)); + + Expr byte_val = bytes(byte_abs, blk); + Expr p3 = mux(n, {1, 3, 9, 27, 81, 243}); + + Expr q_trunc = cast(widening_mul(byte_val, p3)); + Expr xi = cast((cast(q_trunc) * 3) >> 8); + + Func codes("trit_pack_codes"); + codes(kk, blk) = cast(xi); + return codes; + } +}; + +// --------------------------------------------------------------------------- +// 3b. Derived extra fields (computed from other already-encoded fields, not +// from the original values -- appended before struct-packing). +// --------------------------------------------------------------------------- + +// The two ways a derived "sum of codes" extra field gets appended: +// - ScaledFloat: one sum for the *whole* block, already multiplied by +// scale into a float -- GGML's Q8_1 "s" field (group_size == block_size, +// a single group), letting a paired vec_dot recover sum(dequantized +// values) cheaply from the block's own header instead of re-reducing +// codes itself. +// - RawInt16: one sum *per group* (group_size < block_size, several +// groups), each a plain int32-then-int16 integer sum with no scale +// multiply -- GGML's Q8_K "bsums" field, letting a paired K-quant +// vec_dot recover each 16-element group's sum of raw int8 codes cheaply. +enum class SumMode { ScaledFloat, + RawInt16 }; + +// encode({codes(kk, blk), scale(blk)}) -> {codes, scale, sum}: sum(blk) (no +// group dim) for ScaledFloat, sum(g, blk) for RawInt16 -- see SumMode above. +// decode() discards sum and passes codes/scale through unchanged in both +// modes: it's a redundant, derivable quantity, not needed to reconstruct +// dequantized values, so there's nothing to invert. Arity-changing like the +// grouped FieldSpec fields make_block_layout composes (see FieldSpec/ +// FieldLayout above), but in the *encode* direction instead (2 inputs -> 3 +// outputs; decode then undoes it in the same direction rather than the +// mirror one, since sum isn't invertible into anything -- it's simply +// dropped). +class AppendSums { +public: + AppendSums(int group_size, SumMode mode) + : group_size_(group_size), mode_(mode) { + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + Func codes = inputs[0], scale = inputs[1]; + Var blk("blk"); + + if (mode_ == SumMode::ScaledFloat) { + RDom r(0, group_size_, "r"); + Func sum_i("append_sums_i"); + sum_i(blk) = 0; + sum_i(blk) += cast(codes(r, blk)); + + Func sum_f("append_sums_scaled"); + sum_f(blk) = cast(sum_i(blk)) * scale(blk); + + return {codes, scale, sum_f}; + } + Var g("g"); + RDom rg(0, group_size_, "rg"); + + Func sum_i("append_sums_i"); + sum_i(g, blk) = cast(0); + sum_i(g, blk) += cast(codes(g * group_size_ + rg, blk)); + + Func bsums("append_sums_raw"); + bsums(g, blk) = cast(sum_i(g, blk)); + + return {codes, scale, bsums}; + } + + std::vector decode(const std::vector &encoded) const { + // The sum (encoded[2]) is a redundant derived quantity -- pass + // codes/scale through unchanged. + return {encoded[0], encoded[1]}; + } + +private: + int group_size_; + SumMode mode_; +}; + +// --------------------------------------------------------------------------- +// 3c. Concatenation into one byte buffer. +// --------------------------------------------------------------------------- + +// Concatenates N already-packed, fixed-width byte fields into one +// byte-addressed buffer per block, at fixed offsets -- generalizes the +// per-format "select(byte==0, delta_byte0, byte==1, delta_byte1, +// packed(...))" pattern duplicated in every *_generators.cpp quantize +// function today. +// +// `field_widths[k]` is the k-th field's width in bytes, *in on-disk byte +// order*: `inputs[k]`/`encoded[k]` (encode/decode respectively) must already +// be in that same order. When the rest of a Compose/Parallel chain produces +// fields in some other order (e.g. {codes_bytes, scale_bytes} when the +// on-disk layout is scale-then-codes, as in block_q4_0), reorder into byte +// order with a Permute stage composed just before this one (in encode order) -- +// FieldLayout below is the one call site that needs this, via +// Permute{slots_of(fields_)}. +class StructPack { +public: + explicit StructPack(std::vector field_widths) + : field_widths_(std::move(field_widths)) { + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + Var byte("byte"), blk("blk"); + std::vector offsets = offsets_in_output_order(); + + // select()'s branches are all evaluated unconditionally (not + // short-circuiting), so each field's local index is clamped to its + // own valid range before use -- same idiom as q4_0_generators.cpp. + Expr result = cast(0); + for (int k = (int)field_widths_.size() - 1; k >= 0; k--) { + Expr local = clamp(byte - offsets[k], 0, field_widths_[k] - 1); + Expr in_range = byte >= offsets[k] && byte < offsets[k] + field_widths_[k]; + result = select(in_range, inputs[k](local, blk), result); + } + + Func packed("struct_pack_packed"); + packed(byte, blk) = result; + return {packed}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + Func packed = encoded[0]; + std::vector offsets = offsets_in_output_order(); + Var local("local"), blk("blk"); + + std::vector fields; + fields.reserve(field_widths_.size()); + for (int k = 0; k < (int)field_widths_.size(); k++) { + Func field("struct_pack_field_" + std::to_string(k)); + field(local, blk) = packed(local + offsets[k], blk); + fields.push_back(field); + } + return {fields, {}}; + } + +private: + std::vector offsets_in_output_order() const { + std::vector offsets(field_widths_.size()); + int acc = 0; + for (size_t k = 0; k < field_widths_.size(); k++) { + offsets[k] = acc; + acc += field_widths_[k]; + } + return offsets; + } + + std::vector field_widths_; + std::vector input_index_; +}; + +// A struct-typed replacement for StructPack + Permute + the scale field's +// Fp16Pack, backed by a first-class Halide::Type::Struct instead of hand-summed +// byte offsets. `block_type` is the on-disk block's struct type (e.g. +// block_q4_0's `{fp16 d; uint8 qs[16]}`); the compiler owns the field offsets +// and the total byte size (`block_type.bytes()`), so nothing here computes them. +// +// This leaf is the LAST stage of a scheme's Compose (encode order), at the +// on-disk-byte end, exactly where the old make_block_layout() stack was. It produces the same two +// logical Funcs the symmetric/affine quantize stage consumes, in slot order: +// slot 0: `codes_bytes(local, blk)` -- the raw UInt(8) bytes of the codes +// field, still to be interpreted by the code_pack (nibble/byte/bit +// extraction) composed just inside this leaf. Struct types subsume the +// *layout* of these bytes, not the packing trick that reads sub-byte +// codes out of them. +// slot 1: `scale(blk)` -- the block's scale, read straight out of the typed +// `d` field (Float(16) -> Float(32)); this is what subsumes Fp16Pack's +// manual `lo | (hi<<8)` reassembly + reinterpret. +// Supports exactly one scalar scale field + one UInt(8) array codes field (the +// block_q4_0/block_q8_0 shape). Affine (min) and split-code (q5_0) layouts stay +// on make_block_layout for now. +class StructBlockLayout { +public: + StructBlockLayout(Halide::Type block_type, std::string scale_field, std::string codes_field) + : block_type_(block_type), scale_field_(std::move(scale_field)), codes_field_(std::move(codes_field)) { + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + Func codes_bytes = inputs[0]; // codes_bytes(local, blk), UInt(8) + Func scale = inputs[1]; // scale(blk), Float(32) + Var blk("blk"); + + // Assemble one value per field element, in field declaration order, for + // the flattened pack_struct() form. The compiler places each at its own + // offset -- no hand-rolled shifts/masks/reinterpret in the reverse + // direction the way the old encode path needed. + const StructTypeInfo *info = block_type_.struct_type(); + std::vector vals; + for (const StructField &f : info->fields) { + int extent = f.array_extent.value_or(1); + for (int i = 0; i < extent; i++) { + if (f.name == scale_field_) { + vals.push_back(cast(f.type, scale(blk))); + } else if (f.name == codes_field_) { + vals.push_back(cast(f.type, codes_bytes(i, blk))); + } else { + _halide_internal_error << "StructBlockLayout: unexpected field \"" << f.name << "\"\n"; + } + } + } + + Func packed("struct_block_packed"); + packed(blk) = pack_struct(block_type_, vals); + return {packed}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + Func packed = encoded[0]; // packed(blk), struct-typed + Var local("local"), blk("blk"); + + Func scale("struct_block_scale"); + scale(blk) = cast(field(packed(blk), scale_field_)); + + Func codes_bytes("struct_block_codes_bytes"); + codes_bytes(local, blk) = cast(field(packed(blk), codes_field_)[local]); + + return {codes_bytes, scale}; + } + +private: + Halide::Type block_type_; + std::string scale_field_, codes_field_; +}; + +// Struct-typed layout for the split-code Q5_0/Q5_1 blocks. Keeping qh as a +// typed UInt(32) field lets LowerStructTypes expose one four-byte load to LLVM; +// the byte-buffer form presents four unrelated UInt(8) loads, which cannot be +// coalesced before the byte-indexed expansion-table lookups. `codes_bytes` +// retains the existing logical {qh[4], qs[16]} layout, so the reusable +// combined-bit codec inside this leaf is unchanged. +class Q5StructBlockLayout { +public: + Q5StructBlockLayout(Halide::Type block_type, bool affine) + : block_type_(block_type), affine_(affine) { + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + Func codes_bytes = inputs[0]; + Func scale = inputs[1]; + Func min = affine_ ? inputs[2] : Func(); + Var blk("blk"); + + Expr qh = cast(codes_bytes(0, blk)) | + (cast(codes_bytes(1, blk)) << 8) | + (cast(codes_bytes(2, blk)) << 16) | + (cast(codes_bytes(3, blk)) << 24); + std::vector vals = {cast(scale(blk))}; + if (affine_) { + vals.push_back(cast(min(blk))); + } + vals.push_back(qh); + for (int i = 0; i < 16; i++) { + vals.push_back(codes_bytes(i + 4, blk)); + } + + Func packed("q5_struct_block_packed"); + packed(blk) = pack_struct(block_type_, vals); + return {packed}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + Func packed = encoded[0]; + Var local("local"), blk("blk"); + + Func scale("q5_struct_block_scale"); + scale(blk) = cast(field(packed(blk), "d")); + Func qh("q5_struct_block_qh"); + qh(blk) = field(packed(blk), "qh"); + Func qh_bytes("q5_struct_block_qh_bytes"); + qh_bytes(local, blk) = cast(qh(blk) >> (local * 8)); + Func qs_bytes("q5_struct_block_qs_bytes"); + qs_bytes(local, blk) = cast(field(packed(blk), "qs")[local]); + Func codes_bytes("q5_struct_block_codes_bytes"); + codes_bytes(local, blk) = select(local < 4, + qh_bytes(local, blk), + qs_bytes(local - 4, blk)); + + if (affine_) { + Func min("q5_struct_block_min"); + min(blk) = cast(field(packed(blk), "m")); + return {codes_bytes, scale, min}; + } + return {codes_bytes, scale}; + } + +private: + Halide::Type block_type_; + bool affine_; +}; + +// Declares that `pack` maps one logical field to `parts` on-disk fields (and +// back). A positional Parallel divides its Funcs among its children by their +// signatures, and an undeclared unit is assumed to take and produce one Func. +class FieldGroup { +public: + FieldGroup(Halide::Approximation pack, int parts) + : pack_(std::move(pack)), parts_(parts) { + } + + std::vector encode(const std::vector &inputs, + const Halide::ApproximationPorts &input_ports) const { + return pack_.encode(inputs, input_ports).encoded; + } + + std::vector decode(const std::vector &encoded, + const Halide::ApproximationPorts &input_ports) const { + return pack_.decode(encoded, input_ports).decoded; + } + + Halide::ApproximationSignature signature(const Halide::ApproximationPorts &inputs) const { + if (inputs.size() > 1) { + return Halide::ApproximationSignature::unknown(inputs); + } + Halide::ApproximationPorts in = inputs; + if (in.empty()) { + in.emplace_back("field"); + } + Halide::ApproximationPorts out; + for (int i = 0; i < parts_; i++) { + out.emplace_back(in[0].name + "." + std::to_string(i)); + } + return {in, out}; + } + + std::vector children() const { + return {pack_}; + } + +private: + Halide::Approximation pack_; + int parts_; +}; + +// One field, in ON-DISK byte order, of a struct-packed block layout: an +// on-disk byte width plus which logical "slot" it lands in -- the index it +// occupies in the Func vector immediately after StructPack::decode() (and, +// symmetrically, immediately before StructPack::encode()) -- and how to +// pack/unpack it. This is the single declaration that used to be split three +// ways at every make_*_scheme call site: a StructPack{widths, input_index}, +// a stack of per-slot pack stages whose slot numbers had to be kept in sync +// with StructPack's own indices by hand, and (at the Generator call sites) a +// hand-summed block byte count. FieldSpec/make_block_layout below fold all +// three into one list. +// +// Most fields are their own arity-1 group: `pack` set, `arity` left at its +// default of 1, one FieldSpec per on-disk field. A few fields aren't packed +// independently, though -- make_code_pack's code_bits==5 combined codec's +// {nibble, qh} decode into one `codes` field together (Q5_0/Q5_1's split +// 5-bit code), and IQ4XSScalePack's {scales_h, scales_l} decode into one +// `scale` field together. For a group like that, list every physical on-disk +// field it spans (so StructPack still gets each one's own width/slot), but +// only the *leader* -- conventionally, the one whose pack actually does the +// work -- carries `pack` and `arity` (= how many consecutive slots, +// [[slot, slot+arity), the group spans); every other member of the group +// leaves `pack` undefined (a "this slot is spoken for by an earlier FieldSpec's +// group" marker) and `arity` at its default (unused for non-leaders). +struct FieldSpec { + int slot; + int width_bytes; + Halide::Approximation pack; + int arity = 1; +}; + +// The result of make_block_layout(): the assembled layout Approximation, +// ready to compose after (in encode order) a scheme's lossy quantize stage, +// plus the on-disk block's total byte width -- summed here, once, from the +// same field list every make_*_scheme() used to hand-sum separately at its +// own Generator call site. +struct BlockLayout { + Halide::Approximation layout; + int bytes; +}; + +inline BlockLayout make_block_layout(std::vector fields) { + using namespace Halide; + + int bytes = 0; + std::vector widths, permutation; + std::vector leaders; + widths.reserve(fields.size()); + permutation.reserve(fields.size()); + + for (const FieldSpec &f : fields) { + permutation.push_back(f.slot); + widths.push_back(f.width_bytes); + bytes += f.width_bytes; + if (f.pack.defined()) { + leaders.push_back(&f); + } + } + + std::sort(leaders.begin(), leaders.end(), + [](const FieldSpec *a, const FieldSpec *b) { return a->slot < b->slot; }); + + // One child per group of slots, in slot order (Identity for any slot no + // group covers). A group's pack consumes one already-combined logical Func + // in encode() and expands it into `arity` on-disk fields. + std::vector children; + int next_slot = 0; + for (const FieldSpec *f : leaders) { + for (; next_slot < f->slot; next_slot++) { + children.push_back(Identity{}); + } + children.push_back(f->arity == 1 ? f->pack : Approximation(FieldGroup{f->pack, f->arity})); + next_slot += f->arity; + } + for (; next_slot < (int)fields.size(); next_slot++) { + children.push_back(Identity{}); + } + + // Encode order: pack each slot, permute into byte order, then concatenate + // (StructPack; see its doc comment for why the Permute is needed at all). + return {Compose(Parallel(std::move(children)), Permute{permutation}, StructPack{widths}), bytes}; +} + +// --------------------------------------------------------------------------- +// Shared helpers for the extern-delegated formats. +// --------------------------------------------------------------------------- + +// The extern quantize body, shared by every format whose real quantizer is a +// named GGML extern (see ggml_extern_quantize.cpp): it computes nothing +// itself, just names the *_quantize_via_ggml symbol and returns the whole +// packed byte buffer as a single 2-D uint8 Func. Wrapped as ExternQuantize +// below and used as the *encoder* half of a Halide::TrustedInverse, whose +// decoder half is the Compose that unpacks and dequantizes those bytes. +inline std::vector extern_quantize_blocks(const std::vector &inputs, + const std::string &extern_name) { + using namespace Halide; + Func flat = inputs[0]; + Func blocks(extern_name + "_blocks"); + std::vector args = {flat}; + blocks.define_extern(extern_name, args, UInt(8), 2, NameMangling::C); + return {blocks}; +} + +// Decode the 2-byte little-endian fp16 delta stored at bytes(offset)/ +// bytes(offset+1). Returns the delta as an Expr in `blk`. Same bit twiddling +// as Fp16Pack::decode, written inline here because the grid leaves below read +// their delta out of a raw byte buffer at a fixed offset rather than through a +// composed Fp16Pack stage. +inline Halide::Expr fp16_delta(Halide::Func bytes, int offset, Halide::Var blk) { + using namespace Halide; + return cast(reinterpret(cast(le_uint(bytes, offset, blk, 2)))); +} + +// (grid(idx) >> (j*8)) & 0xff -- the byte-within-grid-entry extraction every +// grid leaf below does, whether the grid buffer's entries are 64-bit +// (iq2s_grid/iq2xs_grid/iq1s_grid) or 32-bit (iq3xxs_grid/iq3s_grid/ +// iq2xxs_grid). Templated on the grid's element type so one function covers +// both widths; `j` indexes the byte within the (8- or 4-byte) grid entry. +template +inline Halide::Expr grid_byte(Halide::Buffer grid, Halide::Expr idx, Halide::Expr j) { + using namespace Halide; + Expr grid_val = grid(idx); + return cast((grid_val >> (cast(j) * 8)) & 0xff); +} + +// select(bit, -1.0f, 1.0f) -- the sign-bit-to-multiplier idiom every grid +// leaf's final dequantize multiply uses. +inline Halide::Expr sign_select(Halide::Expr bit) { + using namespace Halide; + return select(bit, -1.0f, 1.0f); +} + +// The ksigns_iq2xs indirection + bit test shared by IQ3_XXS/IQ2_XS/IQ2_XXS: +// look `sign_idx` up in the 128-entry ksigns table, then test bit `j` of the +// looked-up byte. +inline Halide::Expr ksigns_sign(Halide::Buffer ksigns, Halide::Expr sign_idx, Halide::Expr j) { + using namespace Halide; + Expr signs = ksigns(sign_idx); + return (cast(signs) & (cast(1) << j)) != 0; +} + +// select(is_high, byte >> 4, byte & 0x0f) -- the low/high-nibble-of-a-byte +// idiom used throughout the K-quant scale unpackers and the grid/repack +// leaves alike. +inline Halide::Expr nibble_of(Halide::Expr byte_expr, Halide::Expr is_high) { + using namespace Halide; + return select(is_high, byte_expr >> 4, byte_expr & 0x0f); +} + +// Copy one of iq_grids_data.h's static constant codebook tables into a named +// Halide::Buffer the grid classes below index into. +template +inline Halide::Buffer make_grid_buffer(const T *data, int n, const char *name) { + Halide::Buffer buf(n, name); + for (int i = 0; i < n; i++) { + buf(i) = data[i]; + } + return buf; +} + +template +inline Halide::Buffer make_static_codebook(const int8_t (&values)[N], const char *name) { + return Halide::Buffer(const_cast(values), (int)N, name); +} + +// --------------------------------------------------------------------------- +// 4. Extern-delegated quantize + decode-only dequantize-math leaves. +// --------------------------------------------------------------------------- +// +// The formats below (codebook, K-quant, IQ grid, IQ4_XS) all share one shape: +// their forward map (quantize) is an opaque offline black box -- a per-block +// nearest-codeword search, an iterative error-minimizing scale fit, a +// transcendental scale derivation -- that no composition of Halide Funcs +// reproduces bit-for-bit, so it is delegated to a named GGML extern (see +// ggml_extern_quantize.cpp). Their reverse map (dequantize) IS an ordinary, +// bit-exact composition of invertible primitives. Halide::TrustedInverse +// pairs the two: ExternQuantize (encode()) as the encoder half, a plain +// Compose (decode()) as the decoder half. The leaves here are the pieces of +// that decoder that aren't already covered by the packing/reshape components +// in sections 1-3 -- the codebook lookup and the scale-multiply math. Each is +// decode-only: its encode() is exactly the opaque forward map deferred to the +// extern, so it is never called (it only ever lives inside a TrustedInverse's +// decoder) and traps if it somehow is. + +// The encoder half of every extern-delegated format's TrustedInverse: encode() +// delegates to the named GGML extern (extern_quantize_blocks), producing the +// whole packed byte buffer as one 2-D uint8 Func. decode() is never called -- +// the TrustedInverse's decoder half owns dequantize. +class ExternQuantize { +public: + explicit ExternQuantize(std::string extern_name) + : extern_name_(std::move(extern_name)) { + } + + std::vector encode(const std::vector &inputs) const { + return extern_quantize_blocks(inputs, extern_name_); + } + + std::vector decode(const std::vector &) const { + _halide_user_error << "ExternQuantize::decode is never valid -- it is only " + "the encoder half of a TrustedInverse.\n"; + return {}; + } + +private: + std::string extern_name_; +}; + +// The encoder half for formats that have NO forward map at all (the IQ1/IQ2 +// importance-matrix-only quantizers: GGML exposes no *_quantize_via_ggml for +// them). encode() produces a correctly-shaped `blocks(byte, blk)` uint8 Func +// so that Func::approximate_by() -- which always builds encode() before +// decode() -- can splice the round trip; the value is a placeholder (0), +// because Pipeline::sever() always severs this encode and binds the +// real already-quantized Input in its place, so it is never computed. This is +// what lets a decode-only format still go through the standard +// approximate_by/sever path (exercising the framework) instead of a +// bespoke direct-decode generator. decode() traps -- the paired decoder half +// of the TrustedInverse owns dequantize. +class SeveredEncode { +public: + // `dims` is the dimensionality of the packed buffer this stands in for: 2 + // for a plain (byte, blk) codec, 3 for a repack weight buffer + // (byte, k-block, col-group). + explicit SeveredEncode(int block_bytes, int dims = 2) + : block_bytes_(block_bytes), dims_(dims) { + } + + std::vector encode(const std::vector &) const { + using namespace Halide; + std::vector args; + for (int d = 0; d < dims_; d++) { + args.push_back(Var("se" + std::to_string(d))); + } + Func blocks("severed_encode_blocks"); + blocks(args) = cast(0); + return {blocks}; + } + + std::vector decode(const std::vector &) const { + _halide_user_error << "SeveredEncode::decode is never valid -- it is only " + "the (always-severed) encoder half of a TrustedInverse.\n"; + return {}; + } + +private: + int block_bytes_, dims_; +}; + +// codes(kk, blk) -> table[codes], a fixed int8 codebook lookup -- the shared +// codes->value step of every codebook-quantized format (IQ4_NL, MXFP4, TQ1_0, +// TQ2_0, NVFP4, IQ4_XS). Applied to the codes field (in a Parallel) between +// unpacking and the scale multiply, so LinearDequant sees the looked-up value +// in place of a raw integer code. `table` is a Buffer over `static const` +// backing data (matching every per-format lookup_*() helper's idiom), copied +// around as a lightweight handle. encode() is the nearest-codeword search +// deferred to the extern, so it never runs. +class Codebook { +public: + explicit Codebook(Halide::Buffer table) + : table_(table) { + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "Codebook::encode is never valid -- the forward " + "codeword search is deferred to an ExternQuantize.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &codes) const { + using namespace Halide; + Var kk("kk"), blk("blk"); + Buffer table = table_; + // Clamp the index to the table's own extent -- a no-op on valid codes + // (unpackers produce in-range indices), but it gives Halide's bounds + // inference a provable index range across this Func boundary, instead + // of falling back to the full int8 range (accessing table at -128). + // The grid leaves clamp their grid index the same way. + // Dimension-general (pure) decode: the Halide::_ placeholder carries + // any extra trailing "lane" dims (e.g. a matmul weight's column dims) + // through untouched; with zero trailing dims it collapses to the + // familiar (kk, blk). See the DESIGN NOTE by Reblock. + Func values("codebook_values"); + values(kk, blk, _) = table(clamp(cast(codes(kk, blk, _)), 0, table.dim(0).extent() - 1)); + return values; + } + +private: + Halide::Buffer table_; +}; + +// The one decode-only linear dequantize behind every extern-delegated +// format, unifying what used to be two separate leaves (ScaleDequant and +// TwoLevelScaleDequant): +// +// - One-level (has_super_d = false, always has_min = false): the flat +// codebook formats' cast(codes) * scale. Inputs {codes, scale}. +// `sub_size` selects the scale's indexing: 0 means one scale for the +// whole block, a Func with NO sub dimension (scale(blk, _), the shape +// Fp16Pack/F32Pack/E8M0Pack decode to -- IQ4_NL/MXFP4/TQ*); > 0 means +// one scale per `sub_size`-element sub-block, indexed +// scale(kk / sub_size, blk, _) (NVFP4's per-sub-block UE4M3 bytes). +// This is SymmetricAffineQuantize::decode generalized with a sub-block +// scale index; the native symmetric formats keep their own invertible +// SymmetricAffineQuantize, so this is decode-only. +// +// - Two-level (has_super_d = true): the K-quant / IQ4_XS dequantize -- a +// super-block-wide float `d` (and, for the affine K-quants, `dmin`) +// times a per-sub-block scale (and min). Inputs, in order: +// has_min: {d, dmin, scale_min, codes} +// no min: {d, scale, codes} +// When has_min, scale and min arrive *combined* in one func +// scale_min(plane, sub, ...) -- plane 0 = scale, plane 1 = min -- the +// shape PlanarBitPack's plane-axis mode and K4ScaleMinPack both produce, +// so a single field carries both halves of the affine per-sub-block +// parameters (no separate `min` slot to thread). `codes` may be raw +// integer codes (K-quants) or already-looked-up codebook values +// (IQ4_XS, via a Codebook stage). +// +// The two branches keep their exact original float expression shapes +// (multiplication order matters for bit-exactness against GGML's reference +// dequantizers) -- this class only merges the leaves, not the arithmetic. +// Dimension-general via Halide::_: any trailing "lane" dims (a matmul +// weight's columns) ride through untouched; zero of them collapses to the +// familiar (kk, blk). See the DESIGN NOTE by Reblock. +class LinearDequant { +public: + LinearDequant(int sub_size, bool has_super_d, bool has_min) + : sub_size_(sub_size), has_super_d_(has_super_d), has_min_(has_min) { + _halide_user_assert(has_super_d || !has_min) + << "LinearDequant: has_min requires has_super_d (no one-level affine format exists).\n"; + _halide_user_assert(!has_super_d || sub_size > 0) + << "LinearDequant: two-level mode always has a per-sub-block scale.\n"; + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "LinearDequant::encode is never valid -- the forward " + "quantize is deferred to an ExternQuantize.\n"; + return {}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + Var kk("kk"), blk("blk"); + Func dequantized("linear_dequantized"); + if (!has_super_d_) { + Func codes = encoded[0], scale = encoded[1]; + // cast straight off the int8 codes (no int32 detour): the + // extra cast would break change_type()'s + // cast(int8)*cast(int8) SDOT pattern, forcing the vec_dot/ + // repack inner product into a float16 multiply instead of an int8 + // dot. See sdot_schedule.h. + if (sub_size_ == 0) { + dequantized(kk, blk, _) = cast(codes(kk, blk, _)) * scale(blk, _); + } else { + dequantized(kk, blk, _) = cast(codes(kk, blk, _)) * scale(kk / sub_size_, blk, _); + } + } else if (has_min_) { + Func d = encoded[0], dmin = encoded[1], scale_min = encoded[2], codes = encoded[3]; + Expr sub = kk / sub_size_; + dequantized(kk, blk, _) = d(blk, _) * cast(scale_min(0, sub, blk, _)) * cast(codes(kk, blk, _)) - + dmin(blk, _) * cast(scale_min(1, sub, blk, _)); + } else { + Func d = encoded[0], scale = encoded[1], codes = encoded[2]; + dequantized(kk, blk, _) = d(blk, _) * cast(scale(kk / sub_size_, blk, _)) * cast(codes(kk, blk, _)); + } + return {dequantized}; + } + +private: + int sub_size_; + bool has_super_d_, has_min_; +}; + +// --------------------------------------------------------------------------- +// 5. K-quants: combined-bit codes and per-sub-block (scale, min) packing. +// --------------------------------------------------------------------------- + +// The invertible arithmetic behind GGML's K-quant "combined bit" codes: a +// wider-than-one-field code split into a low part and a high part, where +// `code = low + high*high_weight - offset` (verified by hand to collapse +// Q3_K's/Q5_K's/Q6_K's actual bit-OR reconstruction into one formula -- OR +// and + agree because the low/high bit ranges never overlap: high_weight is +// always the low part's own value range). Unlike the extern-delegated leaves +// in section 4, this is genuinely invertible in both directions, so it +// composes as an ordinary symmetric stage: decode() combines {low, high} -> +// code, encode() splits code -> {low, high}. The per-field packing (each part +// through its own PlanarBitPack/BytePack, then StructPack concatenating them +// in the format's on-disk order) is the composition around it -- e.g. for +// Q5_K, Compose{CombineBits{...}, Parallel{low_pack, high_pack}, +// StructPack{{qs, qh}, order}} -- not this leaf, which is only the split/combine math. +class CombineBits { +public: + CombineBits(int high_weight, int offset, bool expanded_high = false) + : high_weight_(high_weight), offset_(offset), expanded_high_(expanded_high) { + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + Func codes = inputs[0]; // codes(kk, blk), the combined (pre-split) value + // `combined` is always >= 0 by construction (offset_ is exactly what + // decode() subtracts after reconstructing low + high*weight), so the + // %// below don't need to handle negative operands. + Expr combined = cast(codes(kk, blk)) + offset_; + Func low("combine_bits_low"); + low(kk, blk) = cast(combined % high_weight_); + Func high("combine_bits_high"); + high(kk, blk) = cast(expanded_high_ ? + (combined / high_weight_) * high_weight_ - offset_ : + combined / high_weight_); + return {low, high}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + Func low = encoded[0], high = encoded[1]; + Func code("combine_bits_code"); + code(kk, blk) = cast(expanded_high_ ? + cast(low(kk, blk)) + cast(high(kk, blk)) : + (cast(low(kk, blk)) + high_weight_ * cast(high(kk, blk))) - offset_); + return {code}; + } + +private: + int high_weight_, offset_; + bool expanded_high_; + Halide::Var kk{"kk"}, blk{"blk"}; +}; + +// (Q2_K's per-sub-block nibble-pair scale/min -- low nibble = scale, high +// nibble = min -- is no longer a bespoke leaf: it is exactly PlanarBitPack's +// plane-axis mode, PlanarBitPack{4, 16, 0, /*plane_axis=*/true}, producing +// the same combined (plane, sub) field K4ScaleMinPack does. See make_q2_k_scheme.) + +// decode(bytes(byte_idx, blk), 12 bytes) -> scale_min(plane, sub, blk) for sub +// in [0, 8), plane 0 = scale / plane 1 = min -- GGML's get_scale_min_k4 scheme, +// shared by Q4_K and Q5_K: for sub<4, scale/min are simply the low 6 bits of +// byte[sub]/byte[sub+4]; for sub>=4, each is a 4-bit low part from byte[sub+4] +// combined with a 2-bit high part borrowed from the top 2 bits of an +// earlier byte (byte[sub-4] for scale, byte[sub] for min) -- a +// bit-interleaved packing that fits 8 six-bit values into 6 bytes' worth of +// budget instead of 8 (see q4_k_generators.cpp's original header comment for +// the full derivation). Emits the combined (plane, sub) field LinearDequant +// consumes. Decode-only: Q4_K/Q5_K quantize is extern-delegated. +class K4ScaleMinPack { +public: + std::vector encode(const std::vector &) const { + _halide_user_error << "K4ScaleMinPack is decode-only -- quantize is deferred to an ExternQuantize.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + // bytes(byte_idx, blk[, ...]), byte_idx in [0, 12) + Var plane("plane"), sub("sub"), blk("blk"); + + Expr jj = clamp(sub - 4, 0, 3); + Expr sc = select(sub < 4, + bytes(sub, blk, _) & 0x3f, + cast((bytes(8 + jj, blk, _) & 0x0f) | ((bytes(jj, blk, _) >> 6) << 4))); + Expr m = select(sub < 4, + bytes(sub + 4, blk, _) & 0x3f, + cast((bytes(8 + jj, blk, _) >> 4) | ((bytes(4 + jj, blk, _) >> 6) << 4))); + + Func scale_min("k4_scale_min_pack"); + scale_min(plane, sub, blk, _) = cast(select(plane == 0, sc, m)); + return scale_min; + } +}; + +// decode(bytes(byte_idx, blk), 12 bytes) -> scale(sub, blk) for sub in +// [0, 16) -- Q3_K's 16 SIGNED 6-bit scale values (no min field), a +// different bit-interleaving than get_scale_min_k4 above: the 2 high bits +// always live in byte (sub%4)+8, at bit-shift 2*(sub/4); the 4 low bits +// live in byte (sub%8), taken from the byte's low nibble if sub<8 or high +// nibble if sub>=8. The final signed value is (low|(high<<4)) - 32 (see +// q3_k_generators.cpp's original header comment for the full derivation). +// Decode-only: Q3_K quantize is extern-delegated. +class Q3KScalePack { +public: + std::vector encode(const std::vector &) const { + _halide_user_error << "Q3KScalePack is decode-only -- quantize is deferred to an ExternQuantize.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + // bytes(byte_idx, blk[, ...]), byte_idx in [0, 12) + Var sub("sub"), blk("blk"); + + // Dimension-general via Halide::_, matching the other scale unpackers. + Expr low_byte_idx = sub % 8; + Expr use_high_nibble = sub >= 8; + Expr low_byte = bytes(low_byte_idx, blk, _); + Expr low_val = cast(nibble_of(low_byte, use_high_nibble)); + Expr high_byte_idx = (sub % 4) + 8; + Expr high_shift = (sub / 4) * 2; + Expr high = cast((bytes(high_byte_idx, blk, _) >> high_shift) & 0x3); + + Func scale("q3k_scale_pack_scale"); + scale(sub, blk, _) = cast((low_val | (high << 4)) - 32); + return scale; + } +}; + +// --------------------------------------------------------------------------- +// 6. IQ2/IQ3 grid+sign codebook dequantize. +// --------------------------------------------------------------------------- +// +// Unlike IQ4_NL/MXFP4/TQ1_0/TQ2_0/NVFP4 above, these codebooks map one index +// to a whole *group* of 4 or 8 signed output bytes at once (GGML's published +// iq2s_grid/iq3xxs_grid/iq3s_grid tables, embedded verbatim from +// iq_grids_data.h), and each format combines its grid index, sign bits, and +// per-group scale via its own distinct bit layout -- there's no shared +// sub-formula across formats the way PlanarBitPack's instances turned out to +// be for the K-quants. Rather than force an +// artificial shared abstraction over 3 genuinely different bit layouts, each +// format below is its own small, decode-only Approximation leaf, wrapped by +// its make_*_scheme() factory in a TrustedInverse{ExternQuantize, Compose{..., +// BlockReshape}} (GGML's own reference quantizer for these runs a per-block +// codebook search -- see ggml_extern_quantize.cpp). decode() is a mechanical, +// verified-unchanged transcription of iq2_s_generators.cpp's/ +// iq3_xxs_generators.cpp's/iq3_s_generators.cpp's own (already bit-exact) +// dequantize math, just reading from a `bytes(byte, blk)` Func instead of +// an `Input>` directly, and producing block-indexed values +// (the composed BlockReshape does the flat<->block reshape) instead of a flat +// row itself. encode() traps: quantize is the ExternQuantize's job. + +// IQ2_S: 256-element superblock, 8 groups of 32 elements, grid index = an 8- +// bit qs byte plus 2 extra high bits from a per-group qh byte (1024-entry, +// 64-bit iq2s_grid, 8 output bytes/index); signs stored directly (no +// ksigns_iq2xs indirection); scale a nibble byte array (2 groups/byte) via +// `d*(0.5+nibble)*0.25` -- {fp16 d; qs[32]; signs[32]; qh[8]; scales[8];}, +// 82 bytes. +class IQ2SGridDequantize { +public: + IQ2SGridDequantize() + : grid_(make_grid_buffer(iq_grids::iq2s_grid, 1024, "iq2s_grid")) { + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "IQ2SGridDequantize is decode-only -- quantize is " + "deferred to an ExternQuantize via TrustedInverse.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + // bytes(byte, blk), byte in [0, 82) + // Superblock structure recovered by the composed BlockReshape({8, 4, 8}): + // j (element in [0,8)), l ([0,4)), ib32 (group in [0,8)). + Var j("j"), l("l"), ib32("ib32"), blk("blk"); + + constexpr int kQsOffset = 2; + constexpr int kSignsOffset = kQsOffset + 256 / 8; // 34 + constexpr int kQhOffset = kSignsOffset + 256 / 8; // 66 + constexpr int kScalesOffset = kQhOffset + 256 / 32; // 74 + + Expr qs_l = bytes(kQsOffset + ib32 * 4 + l, blk); + Expr qh_byte = cast(bytes(kQhOffset + ib32, blk)); + Expr extra_bits = mux(l, {(qh_byte << 8) & 0x300, + (qh_byte << 6) & 0x300, + (qh_byte << 4) & 0x300, + (qh_byte << 2) & 0x300}); + Expr grid_idx_raw = cast(qs_l) + extra_bits; + Expr grid_idx = clamp(cast(cast(grid_idx_raw)), 0, 1023); + + Expr signs_byte = bytes(kSignsOffset + ib32 * 4 + l, blk); + + Expr scales_byte = bytes(kScalesOffset + ib32, blk); + Expr nibble = nibble_of(scales_byte, l >= 2); + + Expr d = fp16_delta(bytes, 0, blk); + Expr db = d * (0.5f + cast(nibble)) * 0.25f; + + Expr gbyte = grid_byte(grid_, grid_idx, j); + Expr sign_bit = (cast(signs_byte) & (cast(1) << j)) != 0; + + Func dequantized("iq2s_grid_dequantized"); + dequantized(j, l, ib32, blk) = db * cast(gbyte) * sign_select(sign_bit); + + return dequantized; + } + +private: + Halide::Buffer grid_; +}; + +// IQ3_XXS: 256-element superblock, 8 groups of 32 elements, TWO grid indices +// per l (plain 8-bit qs bytes, no extra bits; 256-entry, 32-bit +// iq3xxs_grid, 4 output bytes/index); signs via the same ksigns_iq2xs +// indirection as IQ2_XXS (a 7-bit sign_idx and a 4-bit scale exponent both +// bit-packed into one little-endian uint32 "aux32" read from the scales- +// and-signs field); scale via `d*(0.5+exp)*0.5` -- {fp16 d; qs[64]; +// scales_and_signs[32];}, 98 bytes. +class IQ3XXSGridDequantize { +public: + IQ3XXSGridDequantize() + : grid_(make_grid_buffer(iq_grids::iq3xxs_grid, 256, "iq3xxs_grid")), + ksigns_(make_grid_buffer(iq_grids::ksigns_iq2xs, 128, "ksigns_iq2xs")) { + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "IQ3XXSGridDequantize is decode-only -- quantize is " + "deferred to an ExternQuantize via TrustedInverse.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + // bytes(byte, blk), byte in [0, 98) + // Superblock structure recovered by the composed BlockReshape({8, 4, 8}): + // j8 (element in [0,8)), l ([0,4)), ib32 (group in [0,8)). + Var j8("j8"), l("l"), ib32("ib32"), blk("blk"); + Expr j4 = j8 % 4; // byte within the 4-byte grid entry + Expr half = j8 / 4; // 0 (grid1) or 1 (grid2) + + constexpr int kQsOffset = 2; + constexpr int kScalesSignsOffset = kQsOffset + 256 / 4; // 66 + + Expr grid_qs_idx = ib32 * 8 + l * 2 + half; + Expr grid_idx = bytes(kQsOffset + grid_qs_idx, blk); + + Expr aux32 = le_u32(bytes, kScalesSignsOffset + ib32 * 4, blk); + + Expr d = fp16_delta(bytes, 0, blk); + Expr db = d * (0.5f + cast(aux32 >> 28)) * 0.5f; + + Expr sign_idx = cast((aux32 >> (cast(l) * 7)) & 127); + Expr sign_bit = ksigns_sign(ksigns_, sign_idx, j8); + Expr gbyte = grid_byte(grid_, grid_idx, j4); + + Func dequantized("iq3xxs_grid_dequantized"); + dequantized(j8, l, ib32, blk) = db * cast(gbyte) * sign_select(sign_bit); + + return dequantized; + } + +private: + Halide::Buffer grid_; + Halide::Buffer ksigns_; +}; + +// IQ3_S: 256-element superblock, 8 groups of 32 elements, TWO grid indices +// per l (an 8-bit qs byte plus 1 extra high bit from a per-group qh byte, +// combined into a 9-bit index; 512-entry, 32-bit iq3s_grid, 4 output +// bytes/index); signs stored directly (no ksigns_iq2xs indirection, unlike +// IQ2_XXS/IQ3_XXS); scale a nibble byte array (one byte per *pair* of +// groups) via `d*(1+2*nibble)` (odd integers 1,3,...,31 -- not the +// "0.5+x*k" formula every other type here uses) -- {fp16 d; qs[64]; qh[8]; +// signs[32]; scales[4];}, 110 bytes. +class IQ3SGridDequantize { +public: + IQ3SGridDequantize() + : grid_(make_grid_buffer(iq_grids::iq3s_grid, 512, "iq3s_grid")) { + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "IQ3SGridDequantize is decode-only -- quantize is " + "deferred to an ExternQuantize via TrustedInverse.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + // bytes(byte, blk), byte in [0, 110) + // Superblock structure recovered by the composed BlockReshape({8, 4, 8}): + // j8 (element in [0,8)), l ([0,4)), grp (group in [0,8)). + Var j8("j8"), l("l"), grp("grp"), blk("blk"); + Expr j4 = j8 % 4; // byte within the 4-byte grid entry + Expr half = j8 / 4; // 0 (first grid index of this l) or 1 (second) + + constexpr int kQsOffset = 2; + constexpr int kQhOffset = kQsOffset + 256 / 4; // 66 + constexpr int kSignsOffset = kQhOffset + 256 / 32; // 74 + constexpr int kScalesOffset = kSignsOffset + 256 / 8; // 106 + + Expr qs_byte = bytes(kQsOffset + grp * 8 + l * 2 + half, blk); + Expr qh_byte = cast(bytes(kQhOffset + grp, blk)); + Expr bit_pos = l * 2 + half; + Expr high_bit = (qh_byte >> cast(bit_pos)) & 1; + Expr grid_idx = clamp(cast(cast(cast(qs_byte) + (high_bit << 8))), 0, 511); + + Expr signs_byte = bytes(kSignsOffset + grp * 4 + l, blk); + Expr sign_bit = (cast(signs_byte) & (cast(1) << j8)) != 0; + + Expr scales_byte = bytes(kScalesOffset + grp / 2, blk); + Expr nibble = nibble_of(scales_byte, (grp % 2) != 0); + + Expr d = fp16_delta(bytes, 0, blk); + Expr db = d * (1.0f + 2.0f * cast(nibble)); + + Expr gbyte = grid_byte(grid_, grid_idx, j4); + + Func dequantized("iq3s_grid_dequantized"); + dequantized(j8, l, grp, blk) = db * cast(gbyte) * sign_select(sign_bit); + + return dequantized; + } + +private: + Halide::Buffer grid_; +}; + +// The four IQ1/IQ2 importance-matrix-only formats have no from_float extern +// (GGML exposes no *_quantize_via_ggml), so their make_*_scheme() below pairs +// this decode leaf with a SeveredEncode via TrustedInverse (a dequantize-only / +// vec_dot-only round trip; the placeholder encode is always severed). Each is a +// verified-unchanged transcription of the matching *_generators.cpp dequantize, +// reading bytes(byte, blk) and emitting the {8,4,8} superblock form +// (j/j-elem, l, group, blk) so the composed BlockReshape does the reshape. + +// IQ2_XS: 74-byte block. qs[32] as 32 little-endian uint16 (4 per group): low +// 9 bits index the 512-entry uint64 iq2xs_grid, top 7 bits index ksigns_iq2xs; +// scale a nibble byte array (2 l's/byte) via d*(0.5+nibble)*0.25. +class IQ2XSGridDequantize { +public: + IQ2XSGridDequantize() + : grid_(make_grid_buffer(iq_grids::iq2xs_grid, 512, "iq2xs_grid")), + ksigns_(make_grid_buffer(iq_grids::ksigns_iq2xs, 128, "ksigns_iq2xs")) { + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "IQ2XSGridDequantize is decode-only.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + Var j("j"), l("l"), ib32("ib32"), blk("blk"); + constexpr int kQsOffset = 2; + constexpr int kScalesOffset = 66; + + Expr qs_idx = ib32 * 4 + l; + Expr qs_val = le_u16([&](int i) { return bytes(kQsOffset + qs_idx * 2 + i, blk); }); + Expr grid_idx = clamp(cast(qs_val & 511), 0, 511); + Expr sign_idx = clamp(cast(qs_val >> 9), 0, 127); + + Expr scales_byte = bytes(kScalesOffset + ib32, blk); + Expr nibble = nibble_of(scales_byte, l >= 2); + Expr db = fp16_delta(bytes, 0, blk) * (0.5f + cast(nibble)) * 0.25f; + + Expr gbyte = grid_byte(grid_, grid_idx, j); + Expr sign_bit = ksigns_sign(ksigns_, sign_idx, j); + + Func dequantized("iq2xs_grid_dequantized"); + dequantized(j, l, ib32, blk) = db * cast(gbyte) * sign_select(sign_bit); + return dequantized; + } + +private: + Halide::Buffer grid_; + Halide::Buffer ksigns_; +}; + +// IQ2_XXS: 66-byte block. Per group, an 8-byte window: bytes 0..3 are 4 grid +// indices (one per l) into the 256-entry uint64 iq2xxs_grid; bytes 4..7 form a +// uint32 aux32 whose top 4 bits are a scale exponent (d*(0.5+exp)*0.25) and +// whose low 28 bits pack four 7-bit ksigns indices. +class IQ2XXSGridDequantize { +public: + IQ2XXSGridDequantize() + : grid_(make_grid_buffer(iq_grids::iq2xxs_grid, 256, "iq2xxs_grid")), + ksigns_(make_grid_buffer(iq_grids::ksigns_iq2xs, 128, "ksigns_iq2xs")) { + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "IQ2XXSGridDequantize is decode-only.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + Var j("j"), l("l"), ib32("ib32"), blk("blk"); + constexpr int kQsOffset = 2; + + Expr grid_idx = clamp(cast(bytes(kQsOffset + ib32 * 8 + l, blk)), 0, 255); + Expr aux32 = le_u32(bytes, kQsOffset + ib32 * 8 + 4, blk); + Expr db = fp16_delta(bytes, 0, blk) * (0.5f + cast(aux32 >> 28)) * 0.25f; + Expr sign_idx = clamp(cast((aux32 >> (cast(l) * 7)) & 127), 0, 127); + + Expr sign_bit = ksigns_sign(ksigns_, sign_idx, j); + Expr gbyte = grid_byte(grid_, grid_idx, j); + + Func dequantized("iq2xxs_grid_dequantized"); + dequantized(j, l, ib32, blk) = db * cast(gbyte) * sign_select(sign_bit); + return dequantized; + } + +private: + Halide::Buffer grid_; + Halide::Buffer ksigns_; +}; + +// IQ1_S: 50-byte block. qs[32] low grid-index bytes (one per l); qh[8] as 8 +// uint16 (one per group): 3 bits/l give the grid index's high 3 bits, bits +// 12..14 a per-group scale (dl = d*(2*s+1)), bit 15 selects a +/-IQ1S_DELTA +// added to every value. iq1s_grid entries are SIGNED bytes used directly. +class IQ1SGridDequantize { +public: + IQ1SGridDequantize() + : grid_(make_grid_buffer(iq_grids::iq1s_grid, 2048, "iq1s_grid")) { + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "IQ1SGridDequantize is decode-only.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + Var j("j"), l("l"), ib("ib"), blk("blk"); + constexpr float kIQ1S_DELTA = 0.125f; + constexpr int kQsOffset = 2; + constexpr int kQhOffset = 34; + + Expr qh_val = le_u16([&](int i) { return bytes(kQhOffset + ib * 2 + i, blk); }); + Expr dl_scale = cast((qh_val >> 12) & 7); + Expr dl = fp16_delta(bytes, 0, blk) * cast(2 * dl_scale + 1); + Expr delta = select((qh_val & 0x8000) != 0, -kIQ1S_DELTA, kIQ1S_DELTA); + + Expr qs_byte = bytes(kQsOffset + ib * 4 + l, blk); + Expr high3 = cast((qh_val >> (cast(l) * 3)) & 7); + Expr grid_idx = clamp(cast(cast(cast(qs_byte) + (high3 << 8))), 0, 2047); + + Expr grid_signed = reinterpret(grid_byte(grid_, grid_idx, j)); + + Func dequantized("iq1s_grid_dequantized"); + dequantized(j, l, ib, blk) = dl * (cast(grid_signed) + delta); + return dequantized; + } + +private: + Halide::Buffer grid_; +}; + +// IQ1_M: 56-byte block, NO separate delta field. qs[32] low index bytes; +// qh[16] (2/group) give a high-3-bit grid extension + a sign-delta bit per l2; +// scales[8] as 4 uint16: the block's shared fp16 d is bit-gathered from the +// top nibble of all 4 words, each word's low 12 bits holding two 3-bit +// per-group scales. Signed codebook + /-IQ1S_DELTA, same as IQ1_S. +class IQ1MGridDequantize { +public: + IQ1MGridDequantize() + : grid_(make_grid_buffer(iq_grids::iq1s_grid, 2048, "iq1s_grid")) { + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "IQ1MGridDequantize is decode-only.\n"; + return {}; + } + + Halide::Func decode(const Halide::Func &bytes) const { + using namespace Halide; + Var j("j"), l2("l2"), ib("ib"), blk("blk"); + constexpr float kIQ1S_DELTA = 0.125f; + constexpr int kQsOffset = 0; + constexpr int kQhOffset = 32; + constexpr int kScalesOffset = 48; + + auto sc_word = [&](int k) -> Expr { + return le_u16([&](int i) { return bytes(kScalesOffset + k * 2 + i, blk); }); + }; + Expr sc0 = sc_word(0), sc1 = sc_word(1), sc2 = sc_word(2), sc3 = sc_word(3); + Expr d_bits = (sc0 >> 12) | ((sc1 >> 8) & 0xf0) | ((sc2 >> 4) & 0xf00) | (sc3 & 0xf000); + Expr d = cast(reinterpret(cast(d_bits))); + + Expr qh_half = l2 / 2; + Expr parity = l2 % 2; + Expr qh_byte = cast(bytes(kQhOffset + ib * 2 + qh_half, blk)); + Expr qs_byte = bytes(kQsOffset + ib * 4 + l2, blk); + // Constant-amount shifts per parity arm (a variable-amount shift into a + // buffer index defeats Halide bounds inference even when masked). + Expr extra_bits = select(parity == 0, (qh_byte << 8) & 0x700, (qh_byte << 4) & 0x700); + Expr grid_idx = clamp(cast(cast(cast(qs_byte) + extra_bits)), 0, 2047); + + Expr delta_mask = select(parity == 0, cast(0x08), cast(0x80)); + Expr delta = select((qh_byte & delta_mask) != 0, -kIQ1S_DELTA, kIQ1S_DELTA); + + Expr sc_idx = ib / 2; + Expr sc_word_val = mux(sc_idx, {sc0, sc1, sc2, sc3}); + Expr shift = (ib % 2) * 6 + qh_half * 3; + Expr scale3 = (sc_word_val >> cast(shift)) & 7; + Expr dl = d * cast(2 * scale3 + 1); + + Expr grid_signed = reinterpret(grid_byte(grid_, grid_idx, j)); + + Func dequantized("iq1m_grid_dequantized"); + dequantized(j, l2, ib, blk) = dl * (cast(grid_signed) + delta); + return dequantized; + } + +private: + Halide::Buffer grid_; +}; + +// decode({scales_h(2 bytes), scales_l(4 bytes)}) -> scale(sub, blk) for sub in +// [0, 8) -- IQ4_XS's per-sub-block 6-bit scale `ls`, already minus its 32 bias +// so it feeds LinearDequant's `d * scale(sub) * value` directly. `ls` +// is 4 low bits from scales_l (2 sub-blocks/byte) plus 2 high bits from +// scales_h (a little-endian uint16, 2 bits/sub-block) -- a two-field +// bit-interleaving, the peer of Q3KScalePack/K4ScaleMinPack but reading two +// separate byte fields. Decode-only: IQ4_XS quantize is extern-delegated. +class IQ4XSScalePack { +public: + std::vector encode(const std::vector &) const { + _halide_user_error << "IQ4XSScalePack::encode is never valid -- IQ4_XS " + "quantize is deferred to an ExternQuantize.\n"; + return {}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + Func scales_h = encoded[0], scales_l = encoded[1]; + Var sub("sub"), blk("blk"); + + // Dimension-general via Halide::_, matching the other scale unpackers. + // ls = (scales_l[sub/2] >> 4*(sub%2)) & 0xf | ((scales_h >> 2*sub) & 3) << 4 + Expr low4 = cast(nibble_of(scales_l(sub / 2, blk, _), (sub % 2) != 0)); + Expr sh = le_u16([&](int i) { return scales_h(i, blk, _); }); + Expr high2 = cast((sh >> cast(sub * 2)) & 3); + Expr ls = low4 | (high2 << 4); + + Func scale("iq4xs_scale"); + scale(sub, blk, _) = cast(ls - 32); + return {scale}; + } +}; + +// --------------------------------------------------------------------------- +// 7. Repack: interleaved multi-row activation layout (block_q8_0x4 / q8_Kx4). +// +// These are repack-specific instances of the deferred general block-relayout +// (see the DESIGN NOTE by Reblock): "block-layout change prior to applying the +// same lossy quantization Approximations". The 4 interleaved rows are folded +// into the block index blk = ib*4 + row, so the existing (kk, blk) +// SymmetricAffineQuantize + BytePack + Fp16Pack run unchanged; only the two +// relayouts here (input row-blocking and output interleave/header assembly) +// are new. n_rows is fixed at 4 (repack's "x4"). +// --------------------------------------------------------------------------- + +// Losslessly re-view a 2-D activation x(col, row in [0,n_rows)) as +// block-indexed block(kk in [0, block_size), blk = ib*n_rows + row), where ib +// is the k-block. The per-row scale then falls out of the quantizer as a +// per-blk scale. `n_rows` is 4 for every current caller (repack's "x4"), but +// isn't hardcoded here -- callers pass it explicitly. +class RepackRowReshape { +public: + RepackRowReshape(int block_size, int n_rows) + : block_size_(block_size), n_rows_(n_rows) { + } + + Halide::Func encode(const Halide::Func &x) const { + using namespace Halide; + // x(col, row) + Var kk("kk"), blk("blk"); + Func block("repack_row_block"); + block(kk, blk) = x((blk / n_rows_) * block_size_ + kk, blk % n_rows_); + return block; + } + + Halide::Func decode(const Halide::Func &block) const { + using namespace Halide; + // block(kk, blk) + Var col("col"), row("row"); + Func x("repack_row_x"); + x(col, row) = block(col % block_size_, (col / block_size_) * n_rows_ + row); + return x; + } + +private: + int block_size_, n_rows_; +}; + +// Assemble one interleaved output block (byte, ib) from the per-(row) packed +// code bytes and scale bytes of the 4 rows blk = ib*4 + row. Header: 4 deltas +// (`delta_bytes` each, row order). Payload: codes interleaved in groups of +// `blocklen` -- payload position jj -> (row = (jj % (4*bl))/bl, +// kk = (jj/(4*bl))*bl + jj%bl), matching GGML's src_id/src_offset. +class RepackInterleavePack { +public: + RepackInterleavePack(int block_size, int blocklen, int delta_bytes, bool with_bsums = false) + : block_size_(block_size), blocklen_(blocklen), delta_bytes_(delta_bytes), with_bsums_(with_bsums) { + } + + std::vector encode(const std::vector &inputs) const { + using namespace Halide; + Func code = inputs[0], scale = inputs[1]; + Var byte("byte"), ib("ib"); + int header = 4 * delta_bytes_; + int payload_end = header + block_size_ * 4; + + Expr row_d = clamp(byte / delta_bytes_, 0, 3); + Expr delta_byte = scale(clamp(byte % delta_bytes_, 0, delta_bytes_ - 1), ib * 4 + row_d); + + Expr jj = clamp(byte - header, 0, block_size_ * 4 - 1); + Expr row_p = (jj % (4 * blocklen_)) / blocklen_; + Expr kk_p = (jj / (4 * blocklen_)) * blocklen_ + (jj % blocklen_); + Expr code_byte = code(kk_p, ib * 4 + row_p); + + Func blocks("repack_interleave_blocks"); + if (!with_bsums_) { + blocks(byte, ib) = select(byte < header, delta_byte, code_byte); + return {blocks}; + } + + // Q8_K's block_q8_Kx4 appends `bsums`: int16 group-sums of the int8 + // codes, scattered across rows/groups by GGML's index_q8_k mapping. + // Reduce over the interleaved payload order (rj), same as GGML. + Var g("g"); + Func bsums("repack_bsums"); + RDom rj(0, block_size_ * 4, "rj"); + Expr rp = (rj % (4 * blocklen_)) / blocklen_; + Expr kp = (rj / (4 * blocklen_)) * blocklen_ + (rj % blocklen_); + Expr qval = reinterpret(code(kp, ib * 4 + rp)); + int shift = blocklen_ == 8 ? 3 : 2; // log2(blocklen) + Expr idx = (((rj & (4 * blocklen_ - 1)) >> shift) << 2) + ((rj >> 8) << 4) + ((rj >> 6) & 3); + bsums(g, ib) = cast(0); + bsums(idx, ib) += cast(qval); + + int nbsum = (block_size_ / 16) * 4; // 64 groups + Expr bsum_rel = clamp(byte - payload_end, 0, nbsum * 2 - 1); + Expr bsum_g = bsum_rel / 2; + Expr bsum_is_lo = (bsum_rel % 2) == 0; + Expr bsum_bits = reinterpret(cast(bsums(bsum_g, ib))); + Expr bsum_byte = cast(select(bsum_is_lo, bsum_bits & 0xff, (bsum_bits >> 8) & 0xff)); + + blocks(byte, ib) = select(byte < header, delta_byte, byte < payload_end, code_byte, bsum_byte); + return {blocks}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + Func blocks = encoded[0]; + Var kk("kk"), blk("blk"), sbyte("sbyte"); + int header = 4 * delta_bytes_; + + // blk = ib*4 + row: ib = blk/4, row = blk%4. + Func code("repack_interleave_code"); + Expr jj = (kk / blocklen_) * (4 * blocklen_) + (blk % 4) * blocklen_ + (kk % blocklen_); + code(kk, blk) = blocks(header + jj, blk / 4); + Func scaleb("repack_interleave_scale"); + scaleb(sbyte, blk) = blocks((blk % 4) * delta_bytes_ + sbyte, blk / 4); + return {code, scaleb}; + } + +private: + int block_size_, blocklen_, delta_bytes_; + bool with_bsums_; +}; + +// block_q8_0x4 codec: block-layout relayout + the same symmetric Q8_0 quantize +// (amax/127, round) as make_symmetric_block_scheme, interleaved by `blocklen`. +inline Halide::Approximation make_q8_0x4_scheme(int blocklen) { + using namespace Halide; + return Compose( + RepackRowReshape{32, /*n_rows=*/4}, + SymmetricAffineQuantize{32, 127, RoundingMode::Nearest, ScaleAnchor::AbsMax}, + Parallel{{"codes", BytePack{}}, // codes -> bytes + {"scale", Fp16Pack{}}}, // scale -> fp16 bytes + RepackInterleavePack{32, blocklen, /*delta_bytes=*/2}); +} + +// block_q8_Kx4 codec: same shape as make_q8_0x4_scheme but a 256-element block, +// a float32 delta per row (F32Pack), the same round-to-even -127/max Q8_K +// quantize as make_q8_k_scheme, and the interleaved bsums field (with_bsums). +inline Halide::Approximation make_q8_kx4_scheme(int blocklen) { + using namespace Halide; + return Compose( + RepackRowReshape{256, /*n_rows=*/4}, + SymmetricAffineQuantize{256, 127, RoundingMode::NearestEvenClampedHigh, + ScaleAnchor::ExtremeSignedValueTwoStep}, + Parallel{{"codes", BytePack{}}, // codes -> bytes + {"scale", F32Pack{}}}, // scale -> float32 bytes + RepackInterleavePack{256, blocklen, /*delta_bytes=*/4, /*with_bsums=*/true}); +} + +// --------------------------------------------------------------------------- +// 8. Repack weight un-interleave (for gemv/gemm): decode the 3-D interleaved +// weight buffer (byte, k-block, col-group) into per-element codes + per-column +// scale *bytes* (the composed Fp16/F32/E8M0 pack turns those into the float +// scale -- no scale decode duplicated in the leaf), carrying the two column +// dims (col-in-group j, col-group x) as explicit trailing dims so the +// dimension-general LinearDequant/Codebook/scale packs (Halide::_) run unchanged +// and the matmul reduces over (kk, blk). Decode-only (the weight is +// pre-quantized; SeveredEncode is the severed encoder half). This is the +// decode twin of RepackInterleavePack, for the four "simple" weight families. +// --------------------------------------------------------------------------- +enum class RepackWeightCode { SignedByte, // Q8_0: whole signed int8 + SignedNibble, // Q4_0: two's-complement 4-bit (repack's XOR-0x8) + RawNibble }; // IQ4_NL/MXFP4: raw 4-bit codebook index + +class UnInterleaveWeight { +public: + UnInterleaveWeight(int n_cols, int blocklen, int block_size, + RepackWeightCode code_kind, ScaleFormat scale_kind) + : n_cols_(n_cols), blocklen_(blocklen), block_size_(block_size), + code_kind_(code_kind), scale_kind_(scale_kind) { + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "UnInterleaveWeight is decode-only (SeveredEncode is the " + "severed encoder half of its TrustedInverse).\n"; + return {}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + Func blocks = encoded[0]; // blocks(byte, k-block, col-group) + Var byte("byte"), kk("kk"), blk("blk"), j("j"), x("x"); + // Per-column scale header width: fp16 = 2, f32 = 4, e8m0 = 1 byte/column. + int scale_stride = scale_width(scale_kind_); + int header = scale_stride * n_cols_; + + Func codes("uninterleave_codes"); + if (code_kind_ == RepackWeightCode::SignedByte) { + Expr qs_idx = (kk / blocklen_) * n_cols_ * blocklen_ + j * blocklen_ + (kk % blocklen_); + codes(kk, blk, j, x) = reinterpret(blocks(header + qs_idx, blk, x)); + } else { + Expr half = kk / (block_size_ / 2); + Expr el = kk % (block_size_ / 2); + Expr qs_idx = (el / blocklen_) * n_cols_ * blocklen_ + j * blocklen_ + (el % blocklen_); + Expr byte_v = blocks(header + qs_idx, blk, x); + Expr nib = cast(nibble_of(byte_v, half != 0)); + Expr code = code_kind_ == RepackWeightCode::SignedNibble ? select(nib < 8, nib, nib - 16) : nib; + codes(kk, blk, j, x) = cast(code); + } + + // Pure addressing: gather column j's raw scale-header bytes; the composed + // Fp16Pack / F32Pack / E8M0Pack (dimension-general) turns them into the + // float scale, exactly as KQuantDeInterleave emits d_bytes for Fp16Pack. + Func scale_bytes("uninterleave_scale_bytes"); + scale_bytes(byte, blk, j, x) = blocks(scale_stride * j + byte, blk, x); + return {codes, scale_bytes}; + } + +private: + int n_cols_, blocklen_, block_size_; + RepackWeightCode code_kind_; + ScaleFormat scale_kind_; +}; + +// Weight-decode scheme for a "simple" repack weight family (Q4_0/Q8_0/IQ4_NL/ +// MXFP4): un-interleave -> [codebook] -> one-level scale. The col dims ride the +// dimension-general LinearDequant/Codebook via Halide::_. SeveredEncode (3-D) +// stands in for the pre-quantized weight buffer, severed by the gemv/gemm +// generator's sever. +inline Halide::Approximation make_repack_weight_scheme( + int n_cols, int blocklen, int block_bytes, RepackWeightCode code_kind, ScaleFormat scale_kind, + Halide::Buffer table = {}, int block_size = 32) { + using namespace Halide; + // Interpret the scale-header bytes UnInterleaveWeight gathers with the same + // packs the plain codecs use -- no per-kind scale decode duplicated here. + return TrustedInverse( + SeveredEncode{block_bytes, 3}, + Compose{LinearDequant{/*sub_size=*/0, /*has_super_d=*/false, /*has_min=*/false}, + Parallel{Choose{code_kind == RepackWeightCode::RawNibble, Codebook{table}, Identity{}}, + make_scale_pack(scale_kind)}, + UnInterleaveWeight{n_cols, blocklen, block_size, code_kind, scale_kind}}); +} + +// K-quant repack weight decode (block_q{4,5,6,2}_Kx8, n_cols=8): a bespoke, +// verified-unchanged transcription of the repack_gemv_generators.cpp +// weight_value helpers -- the interleaved 256-element super-block with a +// two-level scale (and, for Q4_K/Q5_K, get_scale_min_k4 with the column index +// standing in for the sub-block index). Its decode reads the severed 3-D +// weight buffer and emits Wt(kk, blk, j, x) directly (kept as one leaf rather +// than composed from LinearDequant + the scale-min packs, because the +// interleave is intricate enough that a direct transcription is far less +// error-prone; the SeveredEncode encoder half is severed by the matmul +// generator's sever). blocklen is the interleave width (4 or 8). +enum class KQuantWeightFamily { Q4_K, + Q5_K, + Q6_K, + Q2_K }; + +// The genuinely repack-specific half of a K-quant weight decode: the *byte +// addressing* that maps a logical (element kk, sub-block, column j, col-group +// x) back to the interleaved block's scattered qs/qh/ql bytes and its scale +// region. It is the K-quant analog of UnInterleaveWeight -- pure permutation, +// no dequant arithmetic. It emits the same logical field slots the plain +// K-quant decode's StructPack+packs produce, so the shared downstream +// Compose (Fp16Pack on the fp16 headers, then LinearDequant, in +// LOGICAL element order) finishes the job identically: +// +// Q4_K/Q5_K: {d_bytes, dmin_bytes, scale_min, codes} sub_size 32 +// Q6_K: {d_bytes, scale, codes} sub_size 16 (no min) +// Q2_K: {d_bytes, dmin_bytes, scale_min, codes} sub_size 16 +// +// (has-min families emit scale and min combined in one scale_min(plane, sub, ...) +// field, plane 0 = scale / 1 = min -- the same shape K4ScaleMinPack and +// PlanarBitPack's plane-axis mode produce, so LinearDequant consumes it +// identically.) scale/min are produced as VALUES here (not bytes) because their +// bit layouts can't be delegated to the plain packs: get_scale_min_k4 (Q4_K/Q5_K) +// does its bit-math on what is the *column* index j in the repack while the +// sub-block rides a separate axis -- an axis transpose K4ScaleMinPack's +// (sub-first) interface can't express -- and the qs/qh code stream is +// column-interleaved, so PlanarBitPack's contiguous-window assumption doesn't +// hold. The scale index reduces to kk/sub_size in logical order for all four families +// (verified: Q6_K's base_l/base_h and Q2_K's sm_idx both collapse to kk/16 +// since blocklen divides 16), which is exactly what LinearDequant +// consumes, so no re-order is needed downstream. +class KQuantDeInterleave { +public: + KQuantDeInterleave(KQuantWeightFamily family, int blocklen) + : family_(family), blocklen_(blocklen) { + } + + std::vector encode(const std::vector &) const { + _halide_user_error << "KQuantDeInterleave is decode-only (SeveredEncode is the severed encoder).\n"; + return {}; + } + + std::vector decode(const std::vector &encoded) const { + using namespace Halide; + Func b = encoded[0]; // blocks(byte, k-superblock, col-group) + Var byte("byte"), kk("kk"), blk("blk"), sub("sub"), plane("plane"), j("j"), x("x"); + const int bl = blocklen_; + const int nc = 8; + + // has-min families (Q4_K/Q5_K/Q2_K) emit scale and min combined in one + // field scale_min(plane, sub, ...) (plane 0 = scale, 1 = min), the shape + // LinearDequant consumes; Q6_K (no min) emits a plain scale. + Func codes("kq_codes"), scale("kq_scale"), scale_min("kq_scale_min"), d_bytes("kq_d"), dmin_bytes("kq_dmin"); + + if (family_ == KQuantWeightFamily::Q4_K || family_ == KQuantWeightFamily::Q5_K) { + const bool is_q5 = family_ == KQuantWeightFamily::Q5_K; + const int kScalesOffset = 32; + const int kQsOffset = is_q5 ? 384 : 128; + const int kQhOffset = 128; // Q5_K only + + Expr iter = kk / 64, local64 = kk % 64, half64 = local64 / 32, lpos = local64 % 32; + Expr k_inner = lpos / bl, ii = lpos % bl, k = iter * (32 / bl) + k_inner; + Expr qs_byte = b(kQsOffset + k * nc * bl + j * bl + ii, blk, x); + Expr nibble = cast(nibble_of(qs_byte, half64 != 0)); + Expr value = nibble; + if (is_q5) { + Expr qh_byte = b(kQhOffset + k_inner * (bl * nc) + j * bl + ii, blk, x); + Expr h_bit = cast((cast(qh_byte) >> (iter * 2 + half64)) & 1); + value = nibble | (h_bit << 4); + } + codes(kk, blk, j, x) = value; + + // get_scale_min_k4, addressed by window(sub) and bit-position j, + // scale and min combined into one (plane, sub) field. + Expr window = (sub / 4) * 48 + (sub % 4) * 12; + Expr jj = clamp(j - 4, 0, 3); + Expr sc = select(j < 4, b(kScalesOffset + window + j, blk, x) & 0x3f, + cast((b(kScalesOffset + window + 8 + jj, blk, x) & 0x0f) | + ((b(kScalesOffset + window + jj, blk, x) >> 6) << 4))); + Expr mn = select(j < 4, b(kScalesOffset + window + j + 4, blk, x) & 0x3f, + cast((b(kScalesOffset + window + 8 + jj, blk, x) >> 4) | + ((b(kScalesOffset + window + 4 + jj, blk, x) >> 6) << 4))); + scale_min(plane, sub, blk, j, x) = cast(select(plane == 0, sc, mn)); + + d_bytes(byte, blk, j, x) = b(2 * j + byte, blk, x); + dmin_bytes(byte, blk, j, x) = b(16 + 2 * j + byte, blk, x); + return {d_bytes, dmin_bytes, scale_min, codes}; + } else if (family_ == KQuantWeightFamily::Q6_K) { + const int kScalesOffset = 16, kQlOffset = 144, kQhOffset = 1168; + const int blocks_per_half = 64 / bl; + constexpr int kQlSize = (256 * 8) / 2, kQhSize = (256 * 8) / 4; + + Expr local128 = kk % 128, is_high = local128 >= 64, pos64 = local128 % 64, i = pos64 % bl; + Expr base_l = kk - i - select(is_high, 64, 0), base_h = base_l + 64; + Expr k = (base_l / 128) * blocks_per_half + (base_l % 128) / bl; + Expr ql_byte = b(kQlOffset + clamp(k * nc * bl + j * bl + i, 0, kQlSize - 1), blk, x); + Expr nibble = cast(nibble_of(ql_byte, is_high)); + Expr qh_shift = select(is_high, ((base_h % 128) / 32) * 2, ((base_l % 128) / 32) * 2); + Expr qh_idx_l = (base_l / 128) * 32 + ((base_l + i) % 32); + Expr qh_idx_h = (base_h / 128) * 32 + ((base_h + i) % 32); + Expr qh_off_l = clamp((qh_idx_l / bl) * (bl * nc) + j * bl + (qh_idx_l % bl), 0, kQhSize - 1); + Expr qh_off_h = clamp((qh_idx_h / bl) * (bl * nc) + j * bl + (qh_idx_h % bl), 0, kQhSize - 1); + Expr qh_byte = b(kQhOffset + select(is_high, qh_off_h, qh_off_l), blk, x); + Expr hi2 = cast((qh_byte >> qh_shift) & 3); + codes(kk, blk, j, x) = (nibble | (hi2 << 4)) - 32; + + // 16 plain signed int8 scales, sub = kk/16 (see class comment). + scale(sub, blk, j, x) = cast(reinterpret(b(kScalesOffset + sub * nc + j, blk, x))); + d_bytes(byte, blk, j, x) = b(2 * j + byte, blk, x); + return {d_bytes, scale, codes}; + } else { // Q2_K (blocklen fixed 8 upstream) + const int kDminOffset = 16, kScalesOffset = 32, kQsOffset = 160; + Expr half = kk / 128, local = kk % 128, subg = local / 32, rem32 = local % 32; + Expr k = half * 4 + rem32 / bl, i = rem32 % bl; + Expr qs_byte = b(kQsOffset + k * nc * bl + j * bl + i, blk, x); + codes(kk, blk, j, x) = cast((qs_byte >> (subg * 2)) & 3); + + // scale/min nibble pair, sub = kk/16 (see class comment), combined + // into one (plane, sub) field (plane 0 = scale, 1 = min). + Expr sm_idx = (sub / 8) * 64 + ((sub % 8) / 2) * 16 + j * 2 + (sub % 2); + Expr sm_byte = b(kScalesOffset + sm_idx, blk, x); + scale_min(plane, sub, blk, j, x) = cast(nibble_of(sm_byte, plane != 0)); + d_bytes(byte, blk, j, x) = b(2 * j + byte, blk, x); + dmin_bytes(byte, blk, j, x) = b(kDminOffset + 2 * j + byte, blk, x); + return {d_bytes, dmin_bytes, scale_min, codes}; + } + } + +private: + KQuantWeightFamily family_; + int blocklen_; +}; + +// K-quant repack weight scheme: the de-interleave addressing leaf composed +// with the SAME arithmetic pieces the plain K-quant decode uses (Fp16Pack for +// the fp16 headers, then LinearDequant in logical element order), +// paired with a severed 3-D encoder (the weight is pre-quantized; encode is +// severed by sever). No BlockReshape: the weight stays block-indexed +// (kk, blk) with the column dims (j, x) riding through. +inline Halide::Approximation make_kquant_repack_weight_scheme( + KQuantWeightFamily family, int blocklen, int block_bytes) { + using namespace Halide; + const bool has_min = family != KQuantWeightFamily::Q6_K; + const int sub_size = family == KQuantWeightFamily::Q4_K || family == KQuantWeightFamily::Q5_K ? 32 : 16; + // KQuantDeInterleave's outputs are {d_bytes, [dmin_bytes,] scale(_min), codes}. + std::vector headers{Fp16Pack{}}; // d_bytes -> d + if (has_min) { + headers.push_back(Fp16Pack{}); // dmin_bytes -> dmin + } + headers.push_back(Identity{}); // scale (or scale_min) + headers.push_back(Identity{}); // codes + return TrustedInverse( + SeveredEncode{block_bytes, 3}, + Compose{LinearDequant{sub_size, /*has_super_d=*/true, has_min}, + Parallel(std::move(headers)), + KQuantDeInterleave{family, blocklen}}); +} + +// The make_*() factories below each return one owned Halide::Approximation +// (as a Halide::Approximation, the framework's polymorphic +// scheme handle) -- a single leaf, a Compose, or a TrustedInverse, whichever +// the format actually is, never a single-element Compose wrapper. A Compose +// lists its stages in encode order: stage 0 is innermost (closest to the +// original values) and its last stage is outermost (its encoded output is the +// whole thing's result). +// +// Native (in-Halide-quantizable) formats are a plain Compose. Extern-delegated +// formats -- whose forward quantize is an opaque GGML extern -- are a +// TrustedInverse pairing ExternQuantize (encode) with a Compose (decode); see +// Halide::TrustedInverse and section 4 above. + +// The invertible combined-bit split/combine (see CombineBits above), with +// its own per-part packing folded in via make_block_layout: `fields` lists +// the {low, high} on-disk sub-fields (in on-disk order, tagged with their +// logical slots -- 0 for low, 1 for high, so CombineBits sees {low, high} +// regardless of on-disk order). +inline Halide::Approximation make_combined_bit_codec( + int high_weight, int offset, std::vector fields) { + using namespace Halide; + return Compose( + CombineBits{high_weight, offset}, + make_block_layout(std::move(fields)).layout); +} + +struct CodePackField { + Halide::Approximation pack; + int bytes; +}; + +inline CodePackField make_code_pack(int block_size, int code_bits, int qmax) { + using namespace Halide; + if (code_bits == 4) { + return {nibble_pack(block_size, qmax), block_size / 2}; + } + if (code_bits == 5) { + // Q5_0/Q5_1's split 5-bit code (a 4-bit low nibble plus a 5th high + // bit) -- exactly the K-quant combined-bit-code shape make_q5_k_scheme + // uses for its own adjacent {qh; qs} region (see the deleted + // FiveBitPack's comment, above BitPack, for the verified-by-hand + // equivalence): qh's bit `kk` is element kk's own high bit, i.e. + // le_bit_pack()'s pos_count=1 addressing (unlike Q5_K's per-window + // rotating_bit_pack -- Q5_0/Q5_1's qh is one flat bit array over the + // *whole* block, not windowed); the nibble half is the ordinary + // nibble_pack. `qmax` becomes CombineBits' final recentering offset + // (0 for Q5_1's already-unsigned affine codes) rather than a per-part + // qmax, since the parts here are raw, uncentered digits. + return {Compose( + CombineBits{16, qmax, /*expanded_high=*/true}, + make_block_layout( + {FieldSpec{1, 4, le_bit_pack(16, qmax)}, // qh -> folded high bit and offset + FieldSpec{0, block_size / 2, nibble_pack(block_size)}}) // qs -> low nibble + .layout), + block_size / 2 + 4}; + } + if (code_bits == 1) { + return {BitPack(), block_size / 8}; + } + return {BytePack(), block_size}; +} + +// quantize -> pack codes -> pack scale -> concatenate into one byte buffer +// with the scale stored first (matching every GGML block_* struct's +// `{fp16 d; ...qs;}` layout) -> reshape, plus the on-disk block byte count +// (2 + the code field's own width) -- the source of truth every Generator +// switch used to re-derive by hand. `layout` picks BlockReshape's flat-row +// (quantize_row/dequantize_row) vs block-indexed (vec_dot/repack, whose +// Inputs are already block-indexed so the reshape is a lossless passthrough) +// shape -- see the Layout enum above; this one factory now covers what used +// to be a separate make_symmetric_block_codec()/make_symmetric_block_scheme() +// pair. +inline SchemeAndBytes make_symmetric_block_scheme( + int block_size, int qmax, RoundingMode rounding, ScaleAnchor anchor, int code_bits, + Layout layout = Layout::FlatRow, bool struct_layout = false) { + using namespace Halide; + + if (struct_layout && code_bits == 4) { + _halide_user_assert(block_size == 32 && qmax == 8) + << "The core q4_0 layout requires a 32-element, qmax=8 block\n"; + Type block_type = Type::Struct({{"d", Float(16)}, {"qs", UInt(8), 16}}); + return {Compose( + BlockReshape{block_size, layout == Layout::BlockIndexed}, + SymmetricAffineQuantize{block_size, qmax, rounding, anchor}, + Parallel{{"codes", Compose{nibble_offset(8), PlanarFieldPack{4, 16}}}, + {"scale", fp16_storage()}}, + StructLayout{block_type, {"qs", "d"}}), + block_type.bytes(), + block_type}; + } + + if (struct_layout && code_bits == 8) { + _halide_user_assert(block_size == 32 && qmax == 127) + << "The core q8_0 layout requires a 32-element, qmax=127 block\n"; + Type block_type = Type::Struct({{"d", Float(16)}, {"qs", Int(8), 32}}); + return {Compose( + BlockReshape{block_size, layout == Layout::BlockIndexed}, + SymmetricAffineQuantize{block_size, qmax, rounding, anchor}, + Parallel{{"scale", fp16_storage()}}, + StructLayout{block_type, {"qs", "d"}}), + block_type.bytes(), + block_type}; + } + + auto [code_pack, code_bytes] = make_code_pack(block_size, code_bits, qmax); + + if (struct_layout) { + // Compatibility path for Q1_0. Q4_0/Q8_0 use their faithful, fully + // core-composed layouts above; other symmetric formats migrate + // deliberately rather than changing behavior through this fallback. + Type block_type = Type::Struct({{"d", Float(16)}, {"qs", UInt(8), code_bytes}}); + return {Compose( + BlockReshape{block_size, layout == Layout::BlockIndexed}, + SymmetricAffineQuantize{block_size, qmax, rounding, anchor}, + Parallel{{"codes", std::move(code_pack)}}, + StructBlockLayout{block_type, "d", "qs"}), + block_type.bytes(), + block_type}; + } + + BlockLayout bl = make_block_layout( + {FieldSpec{1, 2, Fp16Pack()}, // scale + FieldSpec{0, code_bytes, std::move(code_pack)}}); // codes + return {Compose( + BlockReshape{block_size, layout == Layout::BlockIndexed}, + SymmetricAffineQuantize{block_size, qmax, rounding, anchor}, + std::move(bl.layout)), + bl.bytes}; +} + +// quantize -> pack codes -> pack scale -> pack min -> concatenate (scale, min, +// codes; matching block_q4_1/block_q5_1's `{fp16 d; fp16 m; ...qs;}` layout) +// -> reshape -- the affine (min+scale) counterpart to +// make_symmetric_block_scheme(), used by Q4_1 (code_bits=4, ClampedInt8). Q5_1 +// pairs with make_affine_5bit_block_scheme() below instead, since it also +// needs the qh high-bit field. +inline SchemeAndBytes make_affine_block_scheme( + int block_size, int levels, AffineRounding rounding, int code_bits, Layout layout = Layout::FlatRow) { + using namespace Halide; + CodePackField code = make_code_pack(block_size, code_bits, /*qmax=*/0); + BlockLayout bl = make_block_layout( + {FieldSpec{1, 2, Fp16Pack()}, // scale + FieldSpec{2, 2, Fp16Pack()}, // min + FieldSpec{0, code.bytes, std::move(code.pack)}}); // codes + return {Compose( + BlockReshape{block_size, layout == Layout::BlockIndexed}, + AffineQuantize{block_size, levels, rounding}, + std::move(bl.layout)), + bl.bytes}; +} + +// Symmetric quantize (like make_symmetric_block_scheme()) but 5-bit -- now +// just make_symmetric_block_scheme with code_bits=5, since make_code_pack's +// code_bits==5 case already assembles the {qh; qs} split (matching +// block_q5_0's `{fp16 d; qh[4]; qs[16];}`) via the combined-bit codec. +// `qmax` is always 16 (5-bit signed range [-16, 15]). Kept as its own named +// entry point (rather than collapsed into SchemeKind::Symmetric in +// symmetric_quant_generators.cpp) so Q5_0/Q5_1's CMakeLists GENERATOR_ARGS +// (kind=symmetric_5bit/affine_5bit) don't need to change in lockstep. +inline SchemeAndBytes make_symmetric_5bit_block_scheme(int block_size, int qmax, + Layout layout = Layout::FlatRow) { + using namespace Halide; + _halide_user_assert(block_size == 32) << "The Q5 struct layout requires a 32-element block.\n"; + _halide_user_assert(qmax == 16) << "The q5_0 additive representation requires qmax 16.\n"; + + // Faithful physical declaration: qh remains an array of four bytes in the + // public type. LittleEndianScalarPack's concat_bits decode lets + // LowerStructTypes recover a single unaligned word load from those bytes. + Type block_type = Type::Struct({{"d", Float(16)}, {"qh", UInt(8), 4}, {"qs", UInt(8), 16}}); + + Approximation qh_word = LittleEndianScalarPack{}; + Approximation radix_split = AdditiveRadixSplit(16, 16); + + auto scheme = Compose( + BlockReshape{block_size, layout == Layout::BlockIndexed}, + SymmetricAffineQuantize{block_size, qmax, RoundingMode::TruncateHalfUpWithOffset, + ScaleAnchor::ExtremeSignedValue}, + Parallel{{"codes", radix_split}}, // codes -> {low, high} + Parallel{{"low", PlanarFieldPack{4, 16}}, + {"high", Compose{BinaryAlphabetPack{32, UInt(32), -16, 0}, qh_word}}, + {"scale", fp16_storage()}}, + StructLayout{block_type, {"qs", "qh", "d"}}); + return {std::move(scheme), block_type.bytes(), block_type, false, + radix_split, qh_word}; +} + +// Affine quantize (like make_affine_block_scheme()) but 5-bit -- likewise now +// just make_affine_block_scheme with code_bits=5, matching block_q5_1's +// `{fp16 d; fp16 m; qh[4]; qs[16];}`. `qmax=0` passed to make_code_pack here +// (unlike Q5_0's 16): AffineQuantize's codes are already unsigned [0, +// levels], not centered, so there's no offset to re-apply before splitting +// into nibble+high-bit. +inline SchemeAndBytes make_affine_5bit_block_scheme(int block_size, int levels, + AffineRounding rounding, + Layout layout = Layout::FlatRow) { + using namespace Halide; + _halide_user_assert(block_size == 32) << "The Q5 struct layout requires a 32-element block.\n"; + Type block_type = Type::Struct({{"d", Float(16)}, {"m", Float(16)}, {"qh", UInt(32)}, {"qs", UInt(8), 16}}); + CodePackField code = make_code_pack(block_size, /*code_bits=*/5, /*qmax=*/0); + return {Compose( + BlockReshape{block_size, layout == Layout::BlockIndexed}, + AffineQuantize{block_size, levels, rounding}, + Parallel{std::vector{std::move(code.pack), Identity{}, Identity{}}}, + Q5StructBlockLayout{block_type, /*affine=*/true}), + block_type.bytes(), + block_type}; +} + +// Symmetric byte-packed quantize (like make_symmetric_block_scheme() with +// code_bits=8) plus AppendSums's derived 's' field (SumMode::ScaledFloat), +// matching block_q8_1's `{fp16 d; fp16 s; qs[32];}` -- Q8_1's scheme. Q8_1 is +// activation-only (GGML has no public to_float for it), so there's normally +// no dequantize_row Generator for this scheme's flat-array variant below -- +// but its decode() is still correct and used by any vec_dot pairing against +// Q8_1 as the activation format. AppendSums needs no Parallel wrapper: it +// consumes and produces the *whole* current list (like quantize itself), +// not just one element of it. +inline SchemeAndBytes make_symmetric_byte_sum_block_scheme(int block_size, int qmax, + Layout layout = Layout::FlatRow) { + using namespace Halide; + BlockLayout bl = make_block_layout( + {FieldSpec{1, 2, Fp16Pack()}, // scale + FieldSpec{2, 2, Fp16Pack()}, // sum + FieldSpec{0, block_size, BytePack()}}); // codes + return {Compose( + BlockReshape{block_size, layout == Layout::BlockIndexed}, + SymmetricAffineQuantize{block_size, qmax, RoundingMode::Nearest, ScaleAnchor::AbsMax}, + AppendSums{block_size, SumMode::ScaledFloat}, + std::move(bl.layout)), + bl.bytes, + /*block_type=*/Halide::Type{}, + /*has_block_sums=*/true}; +} + +// Q8_K: activation-only (quantize_row only, matching Q8_1's own situation +// above -- see q8_k_generators.cpp), one 256-element superblock, plain int8 +// codes (BytePack), one float32 (not fp16) scale (F32Pack), and 16 +// per-group int32-then-int16 sums (AppendSums, SumMode::RawInt16) -- {float d; +// qs[256]; bsums[16];}, 292 bytes. RoundingMode::NearestEvenClampedHigh/ +// ScaleAnchor::ExtremeSignedValueTwoStep reproduce GGML's exact +// nearest-int-then-reciprocal-pair quantizer bit-for-bit -- see their own +// comments in SymmetricAffineQuantize above for why the usual +// Nearest/ExtremeSignedValue formulas aren't equivalent here. +inline SchemeAndBytes make_q8_k_scheme(int block_size, int qmax, Layout layout = Layout::FlatRow) { + using namespace Halide; + BlockLayout bl = make_block_layout( + {FieldSpec{1, 4, F32Pack()}, // scale + FieldSpec{0, block_size, BytePack()}, // codes + FieldSpec{2, (block_size / 16) * 2, Int16Pack()}}); // bsums + return {Compose( + BlockReshape{block_size, layout == Layout::BlockIndexed}, + SymmetricAffineQuantize{block_size, qmax, RoundingMode::NearestEvenClampedHigh, + ScaleAnchor::ExtremeSignedValueTwoStep}, + AppendSums{16, SumMode::RawInt16}, + std::move(bl.layout)), + bl.bytes}; +} + +// --------------------------------------------------------------------------- +// Factory helpers for the shared extern-delegated shapes. Each is a plain +// function that assembles a TrustedInverse{ExternQuantize, Compose{...}} out +// of the section-4 leaves -- transparent (it returns exactly the Compose you'd +// write by hand), unlike the bespoke Approximation subclasses this file used +// to have. The canonical composition for each family lives here once; the +// per-format make_*_scheme() below are just its parameters. +// --------------------------------------------------------------------------- + +// Codebook formats (IQ4_NL/MXFP4/TQ2_0/TQ1_0/NVFP4): unpack the code field, +// look it up in `table`, unpack the scale field (`scale_fmt`/`num_scales` +// picking the pack and its width via make_scale_pack/scale_width, rather +// than a (pack, width) pair kept in sync by hand), one-level scale multiply, +// reshape. `scale_first` is the on-disk field order ({scale; codes} vs +// {codes; scale}); StructPack normalizes both to logical {codes, scale}. +inline SchemeAndBytes make_codebook_scheme( + std::string extern_name, int block_size, Halide::Buffer table, + Halide::Approximation code_pack, int code_bytes, + ScaleFormat scale_fmt, int num_scales, bool scale_first, Layout layout = Layout::FlatRow) { + using namespace Halide; + int scale_bytes = scale_width(scale_fmt, num_scales); + std::vector fields; + if (scale_first) { + fields.push_back({1, scale_bytes, make_scale_pack(scale_fmt)}); + fields.push_back({0, code_bytes, std::move(code_pack)}); + } else { + fields.push_back({0, code_bytes, std::move(code_pack)}); + fields.push_back({1, scale_bytes, make_scale_pack(scale_fmt)}); + } + BlockLayout bl = make_block_layout(std::move(fields)); + return {TrustedInverse( + ExternQuantize{std::move(extern_name)}, + Compose{ + BlockReshape{block_size, layout == Layout::BlockIndexed}, + LinearDequant{num_scales == 1 ? 0 : block_size / num_scales, + /*has_super_d=*/false, /*has_min=*/false}, + Parallel{Codebook{std::move(table)}, // codes -> codebook values + Identity{}}, + std::move(bl.layout), + }), + bl.bytes}; +} + +// K-quant formats (Q2_K/Q3_K/Q4_K/Q5_K/Q6_K): `fields` lists every on-disk +// field (d, [dmin,] scale_min, code -- in on-disk order, tagged with their +// logical slots; see FieldSpec) that make_block_layout unpacks, then the +// two-level scale multiply and reshape. This is the old KQuantDequantize +// parameter list -- now assembling a Compose, not a class. `layout` -- +// unlike every other make_*_scheme() here -- comes before `fields` rather than +// trailing with a default. +inline SchemeAndBytes make_k_quant_scheme( + std::string extern_name, int block_size, int sub_size, bool has_min, Layout layout, + std::vector fields) { + using namespace Halide; + BlockLayout bl = make_block_layout(std::move(fields)); + return {TrustedInverse( + ExternQuantize{std::move(extern_name)}, + Compose{BlockReshape{block_size, layout == Layout::BlockIndexed}, + LinearDequant{sub_size, /*has_super_d=*/true, has_min}, + std::move(bl.layout)}), + bl.bytes}; +} + +// IQ grid formats (IQ2_S/IQ3_XXS/IQ3_S): a bespoke grid+sign+scale decode leaf +// that emits values in the superblock's nested structure, with +// BlockReshape(`block_extents`) doing the flat<->block reshape. `block_bytes` +// is the leaf's own hand-verified on-disk block size -- unlike the +// field-table schemes above it can't be derived from a FieldSpec list (the +// grid leaves are deliberately NOT field-table-decomposed; see section 6's +// design note), so it's declared here, once, next to the leaf that owns it, +// and returned in SchemeAndBytes like every other make_*_scheme(). +inline SchemeAndBytes make_grid_scheme( + std::string extern_name, int block_bytes, Halide::Approximation grid_leaf, + std::vector block_extents, Layout layout = Layout::FlatRow) { + using namespace Halide; + return {TrustedInverse( + ExternQuantize{std::move(extern_name)}, + Compose{BlockReshape{std::move(block_extents), layout == Layout::BlockIndexed}, std::move(grid_leaf)}), + block_bytes}; +} + +inline SchemeAndBytes make_severed_grid_scheme( + int block_bytes, Halide::Approximation grid_leaf, + std::vector block_extents, Layout layout = Layout::FlatRow) { + using namespace Halide; + return {TrustedInverse( + SeveredEncode{block_bytes}, + Compose{BlockReshape{std::move(block_extents), layout == Layout::BlockIndexed}, std::move(grid_leaf)}), + block_bytes}; +} + +// IQ4_NL: 32-element blocks, 4-bit codes into a 16-value non-uniform +// codebook, one fp16 scale per block -- {fp16 d; qs[16];}, 18 bytes. Extern +// quantize; decode unpacks {code_bytes, scale_bytes} (ScaleFirst), looks the +// nibbles up in the codebook, applies the one fp16 scale, and reshapes to a +// flat row. +inline SchemeAndBytes make_iq4_nl_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + static const int8_t kValues[16] = {-127, -104, -83, -65, -49, -35, -22, -10, + 1, 13, 25, 38, 53, 69, 89, 113}; + static const Buffer table = make_static_codebook(kValues, "kvalues_iq4nl"); + return make_codebook_scheme("iq4_nl_quantize_via_ggml", 32, table, + nibble_pack(32), 16, + ScaleFormat::Fp16, /*num_scales=*/1, /*scale_first=*/true, layout); +} + +// MXFP4: 32-element blocks, 4-bit codes into the same-shaped 16-value +// codebook as IQ4_NL (different values), one E8M0 (1-byte, power-of-two) +// scale per block -- {e8m0 e; qs[16];}, 17 bytes. +inline SchemeAndBytes make_mxfp4_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + static const int8_t kValues[16] = {0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12}; + static const Buffer table = make_static_codebook(kValues, "kvalues_mxfp4"); + return make_codebook_scheme("mxfp4_quantize_via_ggml", 32, table, + nibble_pack(32), 16, + ScaleFormat::E8M0, /*num_scales=*/1, /*scale_first=*/true, layout); +} + +// TQ2_0: 256-element superblock, 2-bit codes (each in {0,1,2}, meaning +// {-1,0,1}) via crumb_pack(128)'s window-interleaved layout, one +// fp16 scale -- {qs[64]; fp16 d;}, 66 bytes -- qs *before* d, unlike most +// formats here (StructPack's codes-first field order below). +inline SchemeAndBytes make_tq2_0_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + static const int8_t kValues[4] = {-1, 0, 1, 0}; // index 3 is never produced + static const Buffer table = make_static_codebook(kValues, "kvalues_tq2_0"); + return make_codebook_scheme("tq2_0_quantize_via_ggml", 256, table, + crumb_pack(128), 64, + ScaleFormat::Fp16, /*num_scales=*/1, /*scale_first=*/false, layout); +} + +// TQ1_0: 256-element superblock, base-3 codes (each in {0,1,2}, meaning +// {-1,0,1}) via TritPack's 5-trits/byte (+4-trits/byte tail) packing, one +// fp16 scale -- {qs[48]; qh[4]; fp16 d;}, 54 bytes -- qs+qh (combined, 52 +// bytes) *before* d, like TQ2_0. Reuses TQ2_0's exact {-1, 0, 1, unused} +// codebook (TritPack's codes are the same raw 0/1/2 digit either way). +inline SchemeAndBytes make_tq1_0_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + static const int8_t kValues[4] = {-1, 0, 1, 0}; // index 3 is never produced + static const Buffer table = make_static_codebook(kValues, "kvalues_tq1_0"); + return make_codebook_scheme("tq1_0_quantize_via_ggml", 256, table, + TritPack(), 52, + ScaleFormat::Fp16, /*num_scales=*/1, /*scale_first=*/false, layout); +} + +// NVFP4: 64-element block, 4 sub-blocks of 16 elements each, 4-bit codes +// into the same 16-value codebook MXFP4 uses (NVFP4 is MXFP4 with +// finer-grained scales), one UE4M3 scale *per sub-block* via UE4M3Pack -- +// {d[4]; qs[32];}, 36 bytes -- LinearDequant's num_scales=4 (not 1) is what +// makes each sub-block's dequantize use its own scale byte instead of one +// shared scale for the whole 64-element block. +inline SchemeAndBytes make_nvfp4_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + static const int8_t kValues[16] = {0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12}; + static const Buffer table = make_static_codebook(kValues, "kvalues_nvfp4"); + return make_codebook_scheme("nvfp4_quantize_via_ggml", 64, table, + nibble_pack(16), 32, + ScaleFormat::UE4M3, /*num_scales=*/4, /*scale_first=*/true, layout); +} + +// Q4_K: 256-element superblock, 8 sub-blocks of 32 elements each, plain +// 4-bit codes (nibble_pack(64)) and get_scale_min_k4-packed +// (scale, min) pairs (K4ScaleMinPack) -- {fp16 d; fp16 dmin; scales[12]; +// qs[128];}, 144 bytes, fields already in {d, dmin, scale_min, code} logical +// order. Two-level scale (d*scale(sub)*code - dmin*min(sub)). +inline SchemeAndBytes make_q4_k_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + return make_k_quant_scheme( + "q4_k_quantize_via_ggml", 256, 32, /*has_min=*/true, layout, + {FieldSpec{0, 2, Fp16Pack()}, // d + FieldSpec{1, 2, Fp16Pack()}, // dmin + FieldSpec{2, 12, K4ScaleMinPack()}, // scale_min + FieldSpec{3, 128, nibble_pack(64)}}); // codes +} + +// Q5_K: same super-block/sub-block/scale-min shape as Q4_K, but each code +// is 5 bits: a plain 4-bit low nibble (nibble_pack(64)) plus a +// 5th high bit from a separate 32-byte, 8-window rotating-bit array +// (rotating_bit_pack(32)) -- {fp16 d; fp16 dmin; scales[12]; qh[32]; +// qs[128];}, 176 bytes. qh+qs are adjacent in memory, treated as one +// 160-byte combined "code" field, split by an inner make_combined_bit_codec +// (qh before qs on-disk), offset=0 since Q5_K's code is a plain 0..31 +// unsigned magnitude, not recentered. +inline SchemeAndBytes make_q5_k_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + return make_k_quant_scheme( + "q5_k_quantize_via_ggml", 256, 32, /*has_min=*/true, layout, + {FieldSpec{0, 2, Fp16Pack()}, // d + FieldSpec{1, 2, Fp16Pack()}, // dmin + FieldSpec{2, 12, K4ScaleMinPack()}, // scale_min + // Combined 5-bit code: qh (high bit) + qs (low nibble), on-disk + // qh before qs; offset 0 (plain 0..31). + FieldSpec{3, 160, + make_combined_bit_codec( + 16, 0, + {FieldSpec{1, 32, rotating_bit_pack(32)}, // qh -> high bit + FieldSpec{0, 128, nibble_pack(64)}})}}); // qs -> low nibble +} + +// Q2_K: 256-element superblock, 16 sub-blocks of 16 elements each, plain +// 2-bit codes (crumb_pack(128)) and independent per-sub-block +// nibble-pair (scale, min) via PlanarBitPack's plane-axis mode (low nibble = +// scale = plane 0, high nibble = min = plane 1; no bit-interleaving across +// sub-blocks) -- {scales[16]; qs[64]; fp16 d; fp16 dmin;}, 84 bytes, fields +// on-disk in {scale_min, code, d, dmin} order (scale_min/code *before* d/dmin, +// unlike Q4_K/Q5_K), normalized by their own slots back to {d, dmin, +// scale_min, code}. +inline SchemeAndBytes make_q2_k_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + return make_k_quant_scheme( + "q2_k_quantize_via_ggml", 256, 16, /*has_min=*/true, layout, + {FieldSpec{2, 16, PlanarBitPack(4, 16, 0, /*plane_axis=*/true)}, // scales + FieldSpec{3, 64, crumb_pack(128)}, // qs + FieldSpec{0, 2, Fp16Pack()}, // d + FieldSpec{1, 2, Fp16Pack()}}); // dmin +} + +// Q3_K: 256-element superblock, 16 sub-blocks of 16 elements each, no min +// (symmetric, not affine) -- each code is 3 bits: 2 low bits +// (crumb_pack(128)) plus a high bit from a 32-byte, 8-window +// rotating-bit "hmask" array (rotating_bit_pack(32)), recentered by -4 +// (CombineBits offset=4, matching a signed [-4, 3] range); scale is 16 +// SIGNED 6-bit values, its own bit-interleaving distinct from get_scale_min_k4 +// (Q3KScalePack) -- {hmask[32]; qs[64]; scales[12]; fp16 d;}, 110 bytes. +// hmask+qs are adjacent in memory, treated as one 96-byte combined "code" +// field (hmask before qs on-disk); on-disk {code, scale, d} normalizes to +// logical {d, scale, code}. +inline SchemeAndBytes make_q3_k_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + return make_k_quant_scheme( + "q3_k_quantize_via_ggml", 256, 16, /*has_min=*/false, layout, + {// Combined 3-bit code: hmask (high bit) + qs (low 2 bits). + // On-disk hmask before qs; offset 4 (recenters to signed [-4, 3]). + FieldSpec{2, 96, + make_combined_bit_codec( + 4, 4, + {FieldSpec{1, 32, rotating_bit_pack(32)}, // hmask -> high bit + FieldSpec{0, 64, crumb_pack(128)}})}, // qs -> low 2 bits + FieldSpec{1, 12, Q3KScalePack()}, // scales + FieldSpec{0, 2, Fp16Pack()}}); // d +} + +// Q6_K: 256-element superblock, 16 sub-blocks of 16 elements each, no min -- +// each code is 6 bits: a plain 4-bit low nibble over *two* 128-element +// halves (nibble_pack(128)) plus 2 high bits +// (crumb_pack(128)), recentered by -32; scale is 16 plain SIGNED +// int8 values, no bit-interleaving at all (BytePack -- its plain +// reinterpret is exactly what this needs) -- {ql[128]; qh[64]; +// scales[16]; fp16 d;}, 210 bytes. ql+qh are adjacent in memory, treated as +// one 192-byte combined "code" field (ql before qh on-disk). +inline SchemeAndBytes make_q6_k_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + return make_k_quant_scheme( + "q6_k_quantize_via_ggml", 256, 16, /*has_min=*/false, layout, + {// Combined 6-bit code: ql (low nibble) + qh (high 2 bits). + // On-disk ql before qh; offset 32 (recenters to signed [-32, 31]). + FieldSpec{2, 192, + make_combined_bit_codec( + 16, 32, + {FieldSpec{0, 128, nibble_pack(128)}, // ql -> low nibble + FieldSpec{1, 64, crumb_pack(128)}})}, // qh -> high 2 bits + FieldSpec{1, 16, BytePack()}, // scales, 16 signed int8 values + FieldSpec{0, 2, Fp16Pack()}}); // d +} + +// IQ2_S/IQ3_XXS/IQ3_S: see IQ2SGridDequantize/IQ3XXSGridDequantize/ +// IQ3SGridDequantize above for the bit-layout rationale -- each is a bespoke, +// self-contained grid+sign+scale decode leaf (its three bit layouts share no +// sub-formula worth abstracting). Extern quantize; the grid leaf's decode +// produces block-indexed values, and BlockReshape composes the flat<->block +// reshape on top. +inline SchemeAndBytes make_iq2_s_scheme(Layout layout = Layout::FlatRow) { + return make_grid_scheme("iq2_s_quantize_via_ggml", 82, IQ2SGridDequantize(), {8, 4, 8}, layout); +} + +inline SchemeAndBytes make_iq3_xxs_scheme(Layout layout = Layout::FlatRow) { + return make_grid_scheme("iq3_xxs_quantize_via_ggml", 98, IQ3XXSGridDequantize(), {8, 4, 8}, layout); +} + +inline SchemeAndBytes make_iq3_s_scheme(Layout layout = Layout::FlatRow) { + return make_grid_scheme("iq3_s_quantize_via_ggml", 110, IQ3SGridDequantize(), {8, 4, 8}, layout); +} + +// IQ2_XS/IQ2_XXS/IQ1_S/IQ1_M: importance-matrix-only formats with no forward +// map -- SeveredEncode stands in for the (always-severed) encode half so the +// dequantize/vec_dot still go through approximate_by/sever. Block +// bytes: 74 / 66 / 50 / 56. +inline SchemeAndBytes make_iq2_xs_scheme(Layout layout = Layout::FlatRow) { + return make_severed_grid_scheme(74, IQ2XSGridDequantize(), {8, 4, 8}, layout); +} + +inline SchemeAndBytes make_iq2_xxs_scheme(Layout layout = Layout::FlatRow) { + return make_severed_grid_scheme(66, IQ2XXSGridDequantize(), {8, 4, 8}, layout); +} + +inline SchemeAndBytes make_iq1_s_scheme(Layout layout = Layout::FlatRow) { + return make_severed_grid_scheme(50, IQ1SGridDequantize(), {8, 4, 8}, layout); +} + +inline SchemeAndBytes make_iq1_m_scheme(Layout layout = Layout::FlatRow) { + return make_severed_grid_scheme(56, IQ1MGridDequantize(), {8, 4, 8}, layout); +} + +// IQ4_XS: 256-element superblock, 8 sub-blocks of 32 elements, the superblock +// generalization of IQ4_NL's fixed 16-value codebook -- plain 4-bit codes +// (nibble_pack(32)) into the same kvalues_iq4nl table, scaled by +// `d * (ls - 32)`, a two-level scale (no min) whose per-sub-block `ls` is +// bit-interleaved across two byte fields (IQ4XSScalePack) -- {fp16 d; +// scales_h[2]; scales_l[4]; qs[128];}, 136 bytes. +inline SchemeAndBytes make_iq4_xs_scheme(Layout layout = Layout::FlatRow) { + using namespace Halide; + static const int8_t kValues[16] = {-127, -104, -83, -65, -49, -35, -22, -10, + 1, 13, 25, 38, 53, 69, 89, 113}; + static const Buffer table = make_static_codebook(kValues, "kvalues_iq4nl_xs"); + // scales_h leads a 2-slot group together with scales_l: IQ4XSScalePack's + // decode consumes both (scales_h bytes, then scales_l bytes) to recover + // one `scale(sub)` field, the same grouped-field shape make_code_pack's + // code_bits==5 combined codec uses for Q5_0/Q5_1's {nibble, qh} pair. + BlockLayout bl = make_block_layout( + {FieldSpec{0, 2, Fp16Pack()}, // d + FieldSpec{1, 2, IQ4XSScalePack(), /*arity=*/2}, // scales_h (leads the group) + FieldSpec{2, 4, Halide::Approximation{}}, // scales_l (part of the group above) + FieldSpec{3, 128, nibble_pack(32)}}); // qs -> nibbles + return {TrustedInverse( + ExternQuantize{"iq4_xs_quantize_via_ggml"}, + Compose{ + BlockReshape{256, layout == Layout::BlockIndexed}, + LinearDequant{32, /*has_super_d=*/true, /*has_min=*/false}, + Parallel{Identity{}, Identity{}, + Codebook{table}}, // nibbles -> codebook values ({d, scale, qs}) + std::move(bl.layout), + }), + bl.bytes}; +} + +} // namespace ggml_halide diff --git a/apps/ggml/halide/repack_matmul_generator.cpp b/apps/ggml/halide/repack_matmul_generator.cpp new file mode 100644 index 000000000000..33210ba8a4a8 --- /dev/null +++ b/apps/ggml/halide/repack_matmul_generator.cpp @@ -0,0 +1,285 @@ +// Generic, family-driven repack gemv/gemm, the matmul counterpart of the +// repack quantize_mat codecs. Like the vec_dot generators, the weight and +// activation operands are decoded through the Approximation framework +// (approximate_by + sever), and the interleaved weight layout is a +// lossless relayout (UnInterleaveWeight) composed in front of the same lossy +// quant -- the col dims ride the dimension-general LinearDequant/Codebook via +// Halide::_. One generator backs every (family, n_cols, blocklen) gemv library +// (32 hand-rolled kernels -> 2 generic generators). This file covers the four +// "simple" weight families (Q4_0/Q8_0/IQ4_NL/MXFP4); the interleaved K-quant +// weights layer on top later. + +#include "Halide.h" + +#include "quant_components.h" +#include "sdot_schedule.h" + +using namespace Halide; +using namespace ggml_halide; + +namespace { + +enum class WFamily { Q4_0, + Q8_0, + IQ4_NL, + MXFP4, + Q4_K, + Q5_K, + Q6_K, + Q2_K }; + +inline bool is_kquant(WFamily f) { + return f == WFamily::Q4_K || f == WFamily::Q5_K || f == WFamily::Q6_K || f == WFamily::Q2_K; +} + +struct WeightSpec { + Halide::Approximation scheme; + int block_bytes; +}; + +// Weight decode scheme + on-disk block byte width for a simple family. Byte +// widths: fp16-delta families = 2*n_cols header + payload; mxfp4 = n_cols E8M0 +// header. Nibble payload = 16*n_cols, byte payload = 32*n_cols. +WeightSpec weight_spec(WFamily fam, int n_cols, int blocklen) { + static const int8_t kIq4nl[16] = {-127, -104, -83, -65, -49, -35, -22, -10, + 1, 13, 25, 38, 53, 69, 89, 113}; + static const Buffer iq4nl_lut(const_cast(kIq4nl), 16, "kvalues_iq4nl_gemv"); + static const int8_t kMxfp4[16] = {0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12}; + static const Buffer mxfp4_lut(const_cast(kMxfp4), 16, "kvalues_mxfp4_gemv"); + + switch (fam) { + case WFamily::Q4_0: + return {make_repack_weight_scheme(n_cols, blocklen, 18 * n_cols, + RepackWeightCode::SignedNibble, ScaleFormat::Fp16), + 18 * n_cols}; + case WFamily::Q8_0: + return {make_repack_weight_scheme(n_cols, blocklen, 34 * n_cols, + RepackWeightCode::SignedByte, ScaleFormat::Fp16), + 34 * n_cols}; + case WFamily::IQ4_NL: + return {make_repack_weight_scheme(n_cols, blocklen, 18 * n_cols, + RepackWeightCode::RawNibble, ScaleFormat::Fp16, iq4nl_lut), + 18 * n_cols}; + case WFamily::MXFP4: + return {make_repack_weight_scheme(n_cols, blocklen, 17 * n_cols, + RepackWeightCode::RawNibble, ScaleFormat::E8M0, mxfp4_lut), + 17 * n_cols}; + // K-quant weights (n_cols=8 always): a bespoke interleaved decode leaf. + case WFamily::Q4_K: + return {make_kquant_repack_weight_scheme(KQuantWeightFamily::Q4_K, blocklen, 1152), 1152}; + case WFamily::Q5_K: + return {make_kquant_repack_weight_scheme(KQuantWeightFamily::Q5_K, blocklen, 1408), 1408}; + case WFamily::Q6_K: + return {make_kquant_repack_weight_scheme(KQuantWeightFamily::Q6_K, blocklen, 1680), 1680}; + case WFamily::Q2_K: + return {make_kquant_repack_weight_scheme(KQuantWeightFamily::Q2_K, blocklen, 672), 672}; + } + _halide_internal_error << "RepackGemvGenerator: bad family\n"; + return {}; +} + +// gemv: one plain Q8_0 activation row x every column of a repack-interleaved +// weight matrix -> s(col-in-group, col-group). +class RepackGemvGenerator : public Generator { +public: + GeneratorParam family{ + "family", + WFamily::Q4_0, + {{"q4_0", WFamily::Q4_0}, + {"q8_0", WFamily::Q8_0}, + {"iq4_nl", WFamily::IQ4_NL}, + {"mxfp4", WFamily::MXFP4}, + {"q4_k", WFamily::Q4_K}, + {"q5_k", WFamily::Q5_K}, + {"q6_k", WFamily::Q6_K}, + {"q2_k", WFamily::Q2_K}}}; + GeneratorParam n_cols{"n_cols", 4}; + GeneratorParam blocklen{"blocklen", 4}; + + void configure() { + bool kq = is_kquant(family); + int block_size = kq ? 256 : 32; + WeightSpec w = weight_spec(family, n_cols, blocklen); + // Activation: plain Q8_K for K-quant weights, plain Q8_0 otherwise. + auto act = kq ? make_q8_k_scheme(256, 127, Layout::BlockIndexed).scheme : make_symmetric_block_scheme(32, 127, RoundingMode::Nearest, ScaleAnchor::AbsMax, 8, Layout::BlockIndexed).scheme; + int act_bytes = kq ? (4 + 256 + 2 * (256 / 16)) : (2 + 32); + + ImageParam weight_blocks(UInt(8), 3, "weight_blocks"); // (byte, k-block, col-group) + ImageParam act_blocks(UInt(8), 2, "act_blocks"); // plain Q8_0/Q8_K (byte, k-block) + + Var kk("kk"), blk("blk"), j("j"), x("x"); + Func Wt("wt_naive"), Vec("act_naive"); + Wt(kk, blk, j, x) = 0.0f; + Vec(kk, blk) = 0.0f; + + RDom r(0, block_size, 0, weight_blocks.dim(1).extent(), "r"); + Func s("s"); + s(j, x) = 0.0f; + s(j, x) += Wt(r.x, r.y, j, x) * Vec(r.x, r.y); + + ApproximationResult wr = Wt.approximate_by(w.scheme, {s}); + ApproximationResult ar = Vec.approximate_by(act, {s}); + s.update().eager_inline({wr.replacement, ar.replacement}); + + std::vector sever = wr.encoded; + sever.insert(sever.end(), ar.encoded.begin(), ar.encoded.end()); + std::vector bind = {weight_blocks, act_blocks}; + Pipeline({s}).sever(sever, bind); + + for (Func h : wr.intermediates) { + if (h.has_update_definition()) { + h.compute_root(); + } + } + for (Func h : ar.intermediates) { + if (h.has_update_definition()) { + h.compute_root(); + } + } + + // Simple single-scale weight families (Q4_0/Q8_0/IQ4_NL/MXFP4) reduce + // as a scale-free Int(32) dot (SDOT), same as the vec_dot path: the + // per-block weight scale depends only on (block, column) and the + // activation scale only on the block, so both hoist out of the r.x sum. + // K-quant weights carry two-level (super/sub-block) scales that aren't a + // single per-block-invariant factor, so they keep the default schedule. + if (!kq) { + Var u("u"); + Func s_i32 = sdot_partial(s, {{r.y, u}}, {wr, ar})[0]; + s_i32.compute_root().update().atomic().vectorize(r.x, block_size); + } + + weight_blocks.dim(0).set_bounds(0, w.block_bytes); + weight_blocks.dim(1).set_min(0); + weight_blocks.dim(2).set_min(0); + act_blocks.dim(0).set_bounds(0, act_bytes); + act_blocks.dim(1).set_min(0); + s.output_buffer().dim(0).set_bounds(0, n_cols); + s.output_buffer().dim(1).set_min(0); + + add_input(weight_blocks); + add_input(act_blocks); + add_output(s); + } + + void generate() { + } +}; + +// gemm activation: `nr` rows packed 4-at-a-time into the SAME interleaved +// block layout the weight uses -- so it decodes through make_repack_weight_scheme +// with n_cols=4 (row group of 4), the 4 "columns" being the 4 packed rows. +// Q8_0x4 (fp16 scale, 32-block, 136 B) for simple weights; Q8_Kx4 (f32 scale, +// 256-block, 1168 B incl. dropped bsums) for K-quant weights. +struct ActSpec { + Halide::Approximation scheme; + int block_bytes; + int block_size; +}; +ActSpec act_spec(bool kquant, int blocklen) { + if (kquant) { + // block_q8_Kx4 is 1168 B, but the gemm reads only the f32 d[4] header + // (16 B) + interleaved qs (1024 B) = 1040 B; the 128 B of bsums are + // never touched, so the input is bound to 1040, not the full 1168. + return {make_repack_weight_scheme(4, blocklen, 1040, RepackWeightCode::SignedByte, + ScaleFormat::F32, {}, 256), + 1040, 256}; + } + return {make_repack_weight_scheme(4, blocklen, 136, RepackWeightCode::SignedByte, + ScaleFormat::Fp16), + 136, 32}; +} + +// gemm: 4 packed activation rows x every column of a repack-interleaved weight +// matrix -> s(col-in-group j, col-group x, row-in-group m, row-group y). Both +// operands are interleaved codecs decoded through the framework; the only +// difference from gemv is that the activation is interleaved too (4 rows) and +// the output gains the two activation-lane dims. +class RepackGemmGenerator : public Generator { +public: + GeneratorParam family{ + "family", + WFamily::Q4_0, + {{"q4_0", WFamily::Q4_0}, + {"q8_0", WFamily::Q8_0}, + {"iq4_nl", WFamily::IQ4_NL}, + {"mxfp4", WFamily::MXFP4}, + {"q4_k", WFamily::Q4_K}, + {"q5_k", WFamily::Q5_K}, + {"q6_k", WFamily::Q6_K}, + {"q2_k", WFamily::Q2_K}}}; + GeneratorParam n_cols{"n_cols", 4}; + GeneratorParam blocklen{"blocklen", 4}; + + void configure() { + bool kq = is_kquant(family); + int block_size = kq ? 256 : 32; + WeightSpec w = weight_spec(family, n_cols, blocklen); + ActSpec a = act_spec(kq, blocklen); + + ImageParam weight_blocks(UInt(8), 3, "weight_blocks"); // (byte, k-block, col-group) + ImageParam act_blocks(UInt(8), 3, "act_blocks"); // (byte, k-block, row-group) + + Var kk("kk"), blk("blk"), j("j"), x("x"), m("m"), y("y"); + Func Wt("wt_naive"), Act("act_naive"); + Wt(kk, blk, j, x) = 0.0f; + Act(kk, blk, m, y) = 0.0f; + + RDom r(0, block_size, 0, weight_blocks.dim(1).extent(), "r"); + Func s("s"); + s(j, x, m, y) = 0.0f; + s(j, x, m, y) += Wt(r.x, r.y, j, x) * Act(r.x, r.y, m, y); + + ApproximationResult wr = Wt.approximate_by(w.scheme, {s}); + ApproximationResult ar = Act.approximate_by(a.scheme, {s}); + s.update().eager_inline({wr.replacement, ar.replacement}); + + std::vector sever = wr.encoded; + sever.insert(sever.end(), ar.encoded.begin(), ar.encoded.end()); + std::vector bind = {weight_blocks, act_blocks}; + Pipeline({s}).sever(sever, bind); + + for (Func h : wr.intermediates) { + if (h.has_update_definition()) { + h.compute_root(); + } + } + for (Func h : ar.intermediates) { + if (h.has_update_definition()) { + h.compute_root(); + } + } + + // Simple single-scale weight families reduce as a scale-free Int(32) dot + // (SDOT), same as gemv/vec_dot; K-quant's two-level scales keep the + // default schedule. See sdot_schedule.h. + if (!is_kquant(family)) { + Var u("u"); + Func s_i32 = sdot_partial(s, {{r.y, u}}, {wr, ar})[0]; + s_i32.compute_root().update().atomic().vectorize(r.x, block_size); + } + + weight_blocks.dim(0).set_bounds(0, w.block_bytes); + weight_blocks.dim(1).set_min(0); + weight_blocks.dim(2).set_min(0); + act_blocks.dim(0).set_bounds(0, a.block_bytes); + act_blocks.dim(1).set_min(0); + act_blocks.dim(2).set_min(0); + s.output_buffer().dim(0).set_bounds(0, n_cols); + s.output_buffer().dim(1).set_min(0); + s.output_buffer().dim(2).set_bounds(0, 4); + s.output_buffer().dim(3).set_min(0); + + add_input(weight_blocks); + add_input(act_blocks); + add_output(s); + } + + void generate() { + } +}; + +} // namespace + +HALIDE_REGISTER_GENERATOR(RepackGemvGenerator, repack_gemv) +HALIDE_REGISTER_GENERATOR(RepackGemmGenerator, repack_gemm) diff --git a/apps/ggml/halide/repack_quantize_mat_generators.cpp b/apps/ggml/halide/repack_quantize_mat_generators.cpp new file mode 100644 index 000000000000..76101a4ceb2c --- /dev/null +++ b/apps/ggml/halide/repack_quantize_mat_generators.cpp @@ -0,0 +1,121 @@ +// From-scratch Halide reimplementation of GGML's "repack" quantize_mat +// kernels (see src/ggml-cpu/repack.cpp: ggml_quantize_mat_q8_0_4x4_generic / +// ggml_quantize_mat_q8_0_4x8_generic / ggml_quantize_mat_q8_K_4x4_generic / +// ggml_quantize_mat_q8_K_4x8_generic upstream, as of GGML v0.15.3). These +// take 4 contiguous rows of `k` floats (row r at x[r*k .. r*k+k)) and +// interleave them into ONE activation-format block per `k`-sized chunk, +// where 4 per-row values are laid out consecutively (in groups of +// `blck_size_interleave`) instead of one row's worth at a time -- this is +// the packed activation format the corresponding repack_gemv/repack_gemm +// kernels consume. There are only 4 distinct interleavings (2 activation +// formats x 2 interleave widths), reused across every repack weight type +// that shares that (activation, interleave) pair -- see k_repack_entries in +// providers/ggml_provider.cpp and this file's registration in +// halide_provider.cpp. +// +// block_q8_0x4 layout (136 bytes, one per 32-element chunk x 4 rows): +// byte 0-7: 4 fp16 deltas, one per row, in row order +// byte 8-135: 128 signed int8 quants, interleaved in groups of +// `blck_size_interleave` (4 or 8) per row +// +// block_q8_Kx4 layout (1168 bytes, one per 256-element chunk x 4 rows): +// byte 0-15: 4 float32 deltas, one per row, in row order +// byte 16-1039: 1024 signed int8 quants, interleaved the same way +// byte 1040-1167: 64 signed int16 "bsums" (sum of quants in groups of 16, +// scattered across rows/groups by the same index mapping +// GGML uses -- see index_q8_k below) +// +// Q8_0's per-row scale is the same amax/127 symmetric scale as plain Q8_0 +// (see quant_components.h's make_symmetric_block_scheme()), using round() +// (roundf, not round-to-even). Q8_K's per-row scale is the same -127/max +// signed scale as plain Q8_K (see quant_components.h's make_q8_k_scheme()), +// using the same round-to-nearest-even magic-number trick -- but unlike +// plain Q8_K's quantize_row, this repack version has no final MIN(127, ...) +// clamp (safe here since the scale is derived from this exact block's own +// amax, so values can never exceed +-127 already). +// +// This is intentionally unscheduled beyond the minimum Halide requires for +// legality (an update-defined Func can't stay inline) -- scheduling for +// performance is a later step. + +#include "Halide.h" + +#include "quant_components.h" + +using namespace Halide; +using namespace ggml_halide; + +namespace { + +constexpr int kQK8_0 = 32; +constexpr int kBlockBytesQ8_0x4 = 4 * 2 + kQK8_0 * 4; // 136 + +constexpr int kQK_K = 256; +constexpr int kNumGroups = kQK_K / 16; // 16 +constexpr int kBlockBytesQ8_Kx4 = 4 * 4 + kQK_K * 4 + kNumGroups * 4 * 2; // 1168 + +// Shared "quantize_mat" pipeline: a 2-D activation x(col, row in [0,4)) flows +// through the codec `scheme` (block-relayout + Q8 quantize + interleave) via +// the same approximate_by/sever idiom as codec_generator_base.h's +// Quantize direction -- the encode half is adopted as the output block buffer. +struct QuantizeMatPipe { + ImageParam x; + Func blocks_out; +}; +inline QuantizeMatPipe build_quantize_mat(Approximation scheme, int block_bytes) { + ImageParam x(Float(32), 2, "x"); // dim 0: col-within-row (mult. of block), dim 1: row (4) + Var col("col"), row("row"), byte("byte"), ib("ib"); + Func identity("qm_identity"); + identity(col, row) = x(col, row); + + ApproximationResult r = Func(x).approximate_by(scheme, {identity}); + for (Func h : r.intermediates) { + if (h.has_update_definition()) { + h.compute_root(); + } + } + + ImageParam blocks_in(UInt(8), 2, "blocks_in"); + SeverResult q = Pipeline({identity}).sever(r.encoded, {blocks_in}); + + Func blocks_out("blocks"); + blocks_out(byte, ib) = q.offline.outputs()[0](byte, ib); + blocks_out.output_buffer().dim(0).set_bounds(0, block_bytes); + blocks_out.output_buffer().dim(1).set_min(0); + x.dim(0).set_min(0); + x.dim(1).set_bounds(0, 4); + return {x, blocks_out}; +} + +// q8_0_4x4 / q8_0_4x8 differ only by interleave width (blocklen). +template +class Q8_0QuantizeMatGenerator : public Generator> { +public: + void configure() { + QuantizeMatPipe p = build_quantize_mat(make_q8_0x4_scheme(Blocklen), kBlockBytesQ8_0x4); + this->add_input(p.x); + this->add_output(p.blocks_out); + } + void generate() { + } +}; + +// q8_k_4x4 / q8_k_4x8: same shared pipeline, the Q8_K interleaved codec. +template +class Q8_KQuantizeMatGenerator : public Generator> { +public: + void configure() { + QuantizeMatPipe p = build_quantize_mat(make_q8_kx4_scheme(Blocklen), kBlockBytesQ8_Kx4); + this->add_input(p.x); + this->add_output(p.blocks_out); + } + void generate() { + } +}; + +} // namespace + +HALIDE_REGISTER_GENERATOR(Q8_0QuantizeMatGenerator<4>, q8_0_4x4_quantize_mat) +HALIDE_REGISTER_GENERATOR(Q8_0QuantizeMatGenerator<8>, q8_0_4x8_quantize_mat) +HALIDE_REGISTER_GENERATOR(Q8_KQuantizeMatGenerator<4>, q8_k_4x4_quantize_mat) +HALIDE_REGISTER_GENERATOR(Q8_KQuantizeMatGenerator<8>, q8_k_4x8_quantize_mat) diff --git a/apps/ggml/halide/sdot_schedule.h b/apps/ggml/halide/sdot_schedule.h new file mode 100644 index 000000000000..01b42c866cfb --- /dev/null +++ b/apps/ggml/halide/sdot_schedule.h @@ -0,0 +1,93 @@ +#pragma once + +// Shared "make an Approximation-decoded block reduction accumulate as an +// integer dot product" schedule, used by both the vec_dot and repack matmul +// Generators. Given a reduction `acc` of the shape +// acc(...) += decode(Wt)(r.x, r.y, ...) * decode(Vec)(r.x, r.y) +// whose operands' per-block scales are invariant across the within-block +// reduction r.x (but vary with the block index r.y and any output dims), this +// derives the scale-free Int(32) inner dot the caller then schedules. +// +// Why the deep inline: hoist_invariants() needs each operand's per-block scale +// to appear as a top-level factor of the (rfactored) update product. A single +// eager_inline() of an ApproximationResult's .replacement only peels the +// outermost relayout wrapper; the scale*codes product lives deeper, inside the +// dequantizer Func the Approximation combinators build. eager_inline() no-ops +// on any Func not currently directly called and flattens exposed calls left to +// right, so inlining the whole set of inlinable decode intermediates -- one pass per +// possible chain level -- flattens the decode chains of every operand +// regardless of build order, leaving (codes*scale)*... with the scales as +// loop-invariant leaves. See doc: the SDOT investigation on ggml-on-qk. + +#include "Halide.h" + +#include +#include +#include + +namespace ggml_halide { + +// rfactor `acc` preserving `preserved` (typically {{r.y, u}} -- the block +// index), flatten every operand's decode chain into the resulting per-block +// partial, hoist the now-invariant scales out of the remaining reduction, and +// retype the scale-free inner dot to Int(32). Returns that Int(32) Func -- it +// holds the real reduction, so the caller schedules *it* (compute_root, +// vectorize the within-block RVar, etc.). +// `distribute` multiplies the per-block product out before hoisting, for +// formats whose decode carries an offset: an affine weight makes the product +// (d*code + m) * (d_act*act), which has no single scale to hoist. Multiplied out +// it is d*d_act * sum(code*act) + m*d_act * sum(act), and hoist_invariants() +// gives each term its own accumulator -- both with integer bodies, so both reach +// SDOT. That is ggml's own decomposition of the affine formats. +// `keep_out` identifies decode-chain Funcs that must NOT be flattened -- a caller +// that has scheduled one as a materialization boundary (e.g. Q5_x's +// reconstructed `combine_bits_code`, computed once per block so its qh +// byte->bits table read is a contiguous load instead of a per-lane gather). +// can_be_inlined() only checks purity, so a compute_root schedule alone does not +// stop eager_inline from flattening it; excluding it here does. The scale still +// hoists past it, since it stays an opaque r.x-dependent factor of the product. +inline std::vector sdot_partial(Halide::Func &acc, + const std::vector> &preserved, + const std::vector &operands, + bool distribute = false, + const std::vector &keep_out = {}) { + using namespace Halide; + + Func acc_dot = acc.update().rfactor(preserved); + + auto excluded = [&](const Func &f) { + return std::any_of(keep_out.begin(), keep_out.end(), [&](const Func &kept) { + return kept.defined() && kept.function().same_as(f.function()); + }); + }; + + std::vector decode_funcs; + for (const ApproximationResult &op : operands) { + decode_funcs.push_back(op.replacement); + } + for (const ApproximationResult &op : operands) { + for (const Func &h : op.intermediates) { + if (h.function().can_be_inlined() && !excluded(h)) { + decode_funcs.push_back(h); + } + } + } + for (size_t pass = 0; pass < decode_funcs.size(); pass++) { + acc_dot.update().eager_inline(decode_funcs); + } + + if (distribute) { + acc_dot.update().distribute(); + } + + // One accumulator per term -- one for a symmetric weight, two once an affine + // weight's product has been multiplied out. Each is its own Func, so each + // retypes on its own. + std::vector parts; + for (Func &part : acc_dot.update().hoist_invariants()) { + parts.push_back(part.change_type(Int(32))); + } + return parts; +} + +} // namespace ggml_halide diff --git a/apps/ggml/halide/symmetric_quant_generators.cpp b/apps/ggml/halide/symmetric_quant_generators.cpp new file mode 100644 index 000000000000..86fa7e4dfc65 --- /dev/null +++ b/apps/ggml/halide/symmetric_quant_generators.cpp @@ -0,0 +1,123 @@ +// Generic, GeneratorParam-driven quantize/dequantize pair for GGML's legacy +// per-block quantized formats (see quant_components.h for the reusable +// Approximation pieces this assembles). "Q4_0"/"Q4_1"/"Q5_0"/"Q5_1"/"Q8_0"/ +// "Q8_1" are not distinct C++ classes here -- they're just different +// GENERATOR_ARGS instantiations of the same generator template, registered +// in CMakeLists.txt as e.g. q4_0_quantize/q4_0_dequantize. +// +// Quantize and dequantize share every GeneratorParam (the scheme they +// build is identical, just run in opposite directions), so rather than two +// classes each redeclaring the same params, this is one class template +// parameterized on Direction, following +// apps/linear_algebra/src/blas_l1_generators.cpp's AXPYGenerator +// precedent (one generator template, registered multiple times under +// different names/template args). +// +// The whole pipeline -- for *either* direction -- is built once in +// configure(), not generate(): a single Func::approximate_by() + +// Pipeline::sever() call on a genuinely real ImageParam (not a +// placeholder) produces both an "offline" half (the encode/quantize side, +// still depending on that real ImageParam) and an "online" half (the +// decode/dequantize side, reading from whatever ImageParam +// sever() severed it to instead). Each direction just adopts +// whichever half applies to it as its own Input/Output, via +// GeneratorBase::add_input(const ImageParam&)/add_output(const Func&) -- +// new overloads added to Generator.h/.cpp for exactly this use (see there), +// since the stock add_input>()/add_output>() only ever +// mint fresh, undefined ports for generate() to fill in later. generate() +// is therefore an empty stub: by the time it would run, there's nothing +// left to do. +// +// generate() never calls Approximation::encode()/decode() directly -- only +// through Func::approximate_by() and Pipeline::sever(). This +// configure()/generate() body is identical across every *_quant_generators.cpp +// file in this directory, so it lives in codec_generator_base.h's +// CodecGeneratorBase instead of being repeated here -- this +// class only needs to supply its own GeneratorParams and a build_scheme(). + +#include "Halide.h" + +#include "codec_generator_base.h" +#include "quant_components.h" + +using namespace Halide; +using namespace ggml_halide; + +namespace { + +// Which of quant_components.h's make_*_scheme() factories to use -- the one +// axis that can't be reduced to a GeneratorParam value alone, since each +// scheme needs a different subset/arity of the other params below. +enum class SchemeKind { Symmetric, + Affine, + Symmetric5Bit, + Affine5Bit, + SymmetricByteSum, + Q8K }; + +template +class SymmetricCodecGenerator : public CodecGeneratorBase, dir> { +public: + GeneratorParam block_size{"block_size", 32}; + GeneratorParam qmax{"qmax", 127}; + GeneratorParam code_bits{"code_bits", 8}; + GeneratorParam levels{"levels", 15}; + GeneratorParam rounding{ + "rounding", + RoundingMode::Nearest, + {{"nearest", RoundingMode::Nearest}, + {"truncate_half_up_with_offset", RoundingMode::TruncateHalfUpWithOffset}, + {"sign_only", RoundingMode::SignOnly}}}; + GeneratorParam anchor{ + "anchor", + ScaleAnchor::AbsMax, + {{"abs_max", ScaleAnchor::AbsMax}, + {"extreme_signed", ScaleAnchor::ExtremeSignedValue}, + {"mean_abs", ScaleAnchor::MeanAbs}}}; + GeneratorParam affine_rounding{ + "affine_rounding", + AffineRounding::ClampedInt8, + {{"clamped_int8", AffineRounding::ClampedInt8}, + {"unclamped_uint8", AffineRounding::UnclampedUint8}}}; + GeneratorParam kind{ + "kind", + SchemeKind::Symmetric, + {{"symmetric", SchemeKind::Symmetric}, + {"affine", SchemeKind::Affine}, + {"symmetric_5bit", SchemeKind::Symmetric5Bit}, + {"affine_5bit", SchemeKind::Affine5Bit}, + {"symmetric_byte_sum", SchemeKind::SymmetricByteSum}, + {"q8k", SchemeKind::Q8K}}}; + + SchemeAndBytes build_scheme() const { + // switch's controlling expression can't resolve GeneratorParam's + // implicit conversion operators unambiguously -- .value() sidesteps + // that by returning the plain SchemeKind directly. Each make_*_scheme() + // now returns its own block_bytes alongside the scheme (computed from + // the same field list it builds internally), so there's no byte + // arithmetic to duplicate here. + switch (kind.value()) { + case SchemeKind::Symmetric: + // Phase 3 pilot: the symmetric codecs (Q4_0/Q8_0/Q1_0) build their + // on-disk block as a first-class Type::Struct. + return make_symmetric_block_scheme(block_size, qmax, rounding, anchor, code_bits, + Layout::FlatRow, /*struct_layout=*/true); + case SchemeKind::Affine: + return make_affine_block_scheme(block_size, levels, affine_rounding, code_bits); + case SchemeKind::Symmetric5Bit: + return make_symmetric_5bit_block_scheme(block_size, qmax); + case SchemeKind::Affine5Bit: + return make_affine_5bit_block_scheme(block_size, levels, affine_rounding); + case SchemeKind::SymmetricByteSum: + return make_symmetric_byte_sum_block_scheme(block_size, qmax); + case SchemeKind::Q8K: + return make_q8_k_scheme(block_size, qmax); + } + _halide_internal_error << "unreachable SchemeKind\n"; + } +}; + +} // namespace + +HALIDE_REGISTER_GENERATOR(SymmetricCodecGenerator, symmetric_quantize) +HALIDE_REGISTER_GENERATOR(SymmetricCodecGenerator, symmetric_dequantize) diff --git a/apps/ggml/halide/symmetric_vec_dot_generator.cpp b/apps/ggml/halide/symmetric_vec_dot_generator.cpp new file mode 100644 index 000000000000..24f2dca2e11a --- /dev/null +++ b/apps/ggml/halide/symmetric_vec_dot_generator.cpp @@ -0,0 +1,191 @@ +// Generic, family-driven vec_dot for the symmetric/affine per-block formats, +// the vec_dot counterpart of symmetric_quant_generators.cpp's +// SymmetricCodecGenerator. "q4_0_vec_dot"/"q4_1_vec_dot"/... are PARAMS +// instantiations of this one generator (registered in CMakeLists.txt), not +// per-format C++ classes. Weight and activation are both block-indexed codecs +// from quant_components.h; VecDotGeneratorBase splices them via approximate_by/ +// sever (see vec_dot_generator_base.h) -- generate() never calls +// Approximation::encode()/decode() directly. +// +// The weight is one of the symmetric-family kinds (symmetric / affine / +// symmetric_5bit / affine_5bit); the activation is Q8_0 or Q8_1. Single-scale +// symmetric weights x Q8_0 reach an SDOT Int(32) inner dot; affine (+min) and +// mismatched-block pairings (Q1_0 block 128 x Q8_0 block 32) fall back to a +// Float reduction. + +#include "Halide.h" + +#include "quant_components.h" +#include "vec_dot_generator_base.h" + +using namespace Halide; +using namespace ggml_halide; + +namespace { + +enum class WKind { Symmetric, + Affine, + Symmetric5Bit, + Affine5Bit }; +enum class AKind { Q8_0, + Q8_1 }; + +class SymmetricVecDotGenerator : public VecDotGeneratorBase { +public: + GeneratorParam block_size{"block_size", 32}; + + GeneratorParam w_kind{ + "w_kind", + WKind::Symmetric, + {{"symmetric", WKind::Symmetric}, + {"affine", WKind::Affine}, + {"symmetric_5bit", WKind::Symmetric5Bit}, + {"affine_5bit", WKind::Affine5Bit}}}; + GeneratorParam a_kind{ + "a_kind", + AKind::Q8_0, + {{"q8_0", AKind::Q8_0}, + {"q8_1", AKind::Q8_1}}}; + + // symmetric / symmetric_5bit weight params + GeneratorParam w_qmax{"w_qmax", 8}; + GeneratorParam w_code_bits{"w_code_bits", 4}; + GeneratorParam w_rounding{ + "w_rounding", + RoundingMode::TruncateHalfUpWithOffset, + {{"nearest", RoundingMode::Nearest}, + {"truncate_half_up_with_offset", RoundingMode::TruncateHalfUpWithOffset}, + {"sign_only", RoundingMode::SignOnly}}}; + GeneratorParam w_anchor{ + "w_anchor", + ScaleAnchor::ExtremeSignedValue, + {{"abs_max", ScaleAnchor::AbsMax}, + {"extreme_signed", ScaleAnchor::ExtremeSignedValue}, + {"mean_abs", ScaleAnchor::MeanAbs}}}; + + // affine / affine_5bit weight params + GeneratorParam w_levels{"w_levels", 15}; + GeneratorParam w_affine_rounding{ + "w_affine_rounding", + AffineRounding::ClampedInt8, + {{"clamped_int8", AffineRounding::ClampedInt8}, + {"unclamped_uint8", AffineRounding::UnclampedUint8}}}; + + // activation param (Q8_0/Q8_1 are always 8-bit int8 codes) + GeneratorParam a_qmax{"a_qmax", 127}; + + VecDotSpec build_vec_dot() const { + int wbs = block_size; + + Halide::Approximation wc; + int wb; + ScheduleKind sched; + bool distribute = false; // set for the affine (offset-carrying) weights + Halide::Type weight_type; // set -> weight blocks are a 1-D Type::Struct buffer + Halide::Approximation reconstructed_codes_stage, packed_high_word_stage; + switch (w_kind.value()) { + case WKind::Symmetric: { + // Struct-typed weight blocks (`{fp16 d; uint8 qs[...]}`); SDOT still + // works because the base header's deep inline flattens the struct + // decode's dequantizer just like the byte path's. + SchemeAndBytes sb = make_symmetric_block_scheme(wbs, w_qmax, w_rounding, w_anchor, w_code_bits, + Layout::BlockIndexed, /*struct_layout=*/true); + wc = std::move(sb.scheme); + weight_type = sb.block_type; + wb = 2 + (w_code_bits == 4 ? wbs / 2 : (w_code_bits == 1 ? wbs / 8 : wbs)); + // 1-bit (Q1_0) stays Float: change_type(Int(32)) can't prove its + // deep-inlined per-term range fits Int(32) (its 128-wide block trips + // the overflow check), and it's a niche format. Q4_0/Q8_0 -> SDOT. + sched = w_code_bits == 1 ? ScheduleKind::Float : ScheduleKind::SDOT; + break; + } + case WKind::Affine: + wc = make_affine_block_scheme(wbs, w_levels, w_affine_rounding, w_code_bits, Layout::BlockIndexed).scheme; + wb = 2 + 2 + (w_code_bits == 4 ? wbs / 2 : wbs); + // The affine decode is (d*code + m), so the per-block product is not + // a single scaled term. hoist_invariants() multiplies it out into + // d*d_act * sum(code*act) + m*d_act * sum(act) + // and gives each term its own accumulator, which is ggml's own + // decomposition -- both inner sums are integer, so both reach SDOT. + sched = ScheduleKind::SDOT; + distribute = true; + break; + case WKind::Symmetric5Bit: { + SchemeAndBytes sb = make_symmetric_5bit_block_scheme(wbs, w_qmax, Layout::BlockIndexed); + wc = std::move(sb.scheme); + weight_type = sb.block_type; + reconstructed_codes_stage = sb.reconstructed_codes_stage; + packed_high_word_stage = sb.packed_high_word_stage; + wb = 2 + 4 + wbs / 2; + // Q5_0's 5-bit code is assembled via CombineBits (nibble | + // (high_bit << 4)); that reconstruction is all inside the (r.x- + // dependent) codes leaf, so the per-block scale is still a top-level + // invariant factor once the decode chain is fully inlined -- SDOT + // works the same as Q4_0 (see the base header's deep-inline SDOT). + sched = ScheduleKind::SDOT; + break; + } + case WKind::Affine5Bit: { + SchemeAndBytes sb = make_affine_5bit_block_scheme(wbs, w_levels, w_affine_rounding, Layout::BlockIndexed); + wc = std::move(sb.scheme); + weight_type = sb.block_type; + wb = 2 + 2 + 4 + wbs / 2; + // Same offset-carrying decode as the 4-bit affine case above; the + // 5-bit code reconstruction stays inside the codes leaf, so the + // multiplied-out terms are integer just the same. + sched = ScheduleKind::SDOT; + distribute = true; + break; + } + } + + // Q8_0/Q8_1 activations are 32-element blocks. Build the codec at that + // natural block size, then Reblock to the weight's block size (a no-op + // when they already match, e.g. Q4_0/Q8_0); the byte width stays the + // natural-block width since y_blocks is stored at 32-element blocks. + const int a_nat = 32; + Halide::Approximation ac; + int ab; + bool act_has_block_sums = false; + Halide::Type act_type; + switch (a_kind.value()) { + case AKind::Q8_0: { + // Keep the shared activation ABI byte-addressed. A struct-typed + // activation was tested for same-size blocks, but disturbed the + // tuned q4_0/q5_0 load shape; the faithful struct scheme remains in + // use for q8_0's codecs and weight operand. + const bool structured_q8 = false; + SchemeAndBytes sb = make_symmetric_block_scheme( + a_nat, a_qmax, RoundingMode::Nearest, ScaleAnchor::AbsMax, 8, + Layout::BlockIndexed, structured_q8); + ac = std::move(sb.scheme); + act_type = sb.block_type; + ab = 2 + a_nat; + break; + } + case AKind::Q8_1: { + SchemeAndBytes sb = make_symmetric_byte_sum_block_scheme(a_nat, a_qmax, Layout::BlockIndexed); + ac = std::move(sb.scheme); + ab = 2 + 2 + a_nat; + // Q8_1 stores its per-block scaled code sum (`s`); the affine vec_dot + // severs the offset term straight to it. Only meaningful when the + // activation block matches the weight block (a_nat == wbs, i.e. no + // Reblock) -- true for the q4_1/q5_1 pairings. + act_has_block_sums = sb.has_block_sums && a_nat == wbs; + break; + } + } + ac = reblock_activation(std::move(ac), a_nat, wbs); + + const bool five_bit = w_kind.value() == WKind::Symmetric5Bit || w_kind.value() == WKind::Affine5Bit; + const bool clean_symmetric = w_kind.value() == WKind::Symmetric && + (w_code_bits == 4 || w_code_bits == 8); + return {std::move(wc), wb, std::move(ac), ab, wbs, sched, distribute, weight_type, act_type, act_has_block_sums, + five_bit ? 2 : 4, reconstructed_codes_stage, packed_high_word_stage, + clean_symmetric, clean_symmetric}; + } +}; + +} // namespace + +HALIDE_REGISTER_GENERATOR(SymmetricVecDotGenerator, symmetric_vec_dot) diff --git a/apps/ggml/halide/test_bf16.cpp b/apps/ggml/halide/test_bf16.cpp new file mode 100644 index 000000000000..87aa15a2ce74 --- /dev/null +++ b/apps/ggml/halide/test_bf16.cpp @@ -0,0 +1,45 @@ +// Standalone round-trip check: the from-scratch Halide BF16 kernels (both +// directions fully native) vs GGML's own reference implementation, reached +// via the public ggml_get_type_traits() API. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_BF16); + const size_t out_bytes = ggml_row_size(GGML_TYPE_BF16, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_bf16(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_bf16(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_f16.cpp b/apps/ggml/halide/test_f16.cpp new file mode 100644 index 000000000000..ec1f8f619f7b --- /dev/null +++ b/apps/ggml/halide/test_f16.cpp @@ -0,0 +1,45 @@ +// Standalone round-trip check: the from-scratch Halide F16 kernels (both +// directions fully native) vs GGML's own reference implementation, reached +// via the public ggml_get_type_traits() API. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_F16); + const size_t out_bytes = ggml_row_size(GGML_TYPE_F16, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_f16(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_f16(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_iq1_m.cpp b/apps/ggml/halide/test_iq1_m.cpp new file mode 100644 index 000000000000..1f9324bee8d1 --- /dev/null +++ b/apps/ggml/halide/test_iq1_m.cpp @@ -0,0 +1,48 @@ +// Standalone check: the from-scratch Halide IQ1_M dequantize kernel vs +// GGML's own reference implementation. GGML has no public from_float_ref +// for this importance-matrix-only codebook type (see +// ../providers/ggml_internal_abi.h's quantize_iq1_m doc comment) -- its +// only quantizer is reached the same way ggml_provider.cpp reaches it: the +// private whole-matrix quantize_iq1_m symbol, with a uniform (all-1.0) +// weighting for consistency with IQ2_XXS/IQ2_XS/IQ1_S (whose equivalent +// weights are a hard requirement, not just optional here). + +#include +#include +#include +#include + +#include + +#include "../providers/ggml_internal_abi.h" +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + ggml_quantize_init(GGML_TYPE_IQ1_M); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_IQ1_M); + const size_t out_bytes = ggml_row_size(GGML_TYPE_IQ1_M, k); + + std::vector weights(k, 1.0f); + std::vector ref_blocks(out_bytes); + quantize_iq1_m(x.data(), ref_blocks.data(), /*nrows=*/1, /*n_per_row=*/k, weights.data()); + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_iq1_m(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_iq1_s.cpp b/apps/ggml/halide/test_iq1_s.cpp new file mode 100644 index 000000000000..67703e23492a --- /dev/null +++ b/apps/ggml/halide/test_iq1_s.cpp @@ -0,0 +1,47 @@ +// Standalone check: the from-scratch Halide IQ1_S dequantize kernel vs +// GGML's own reference implementation. GGML has no public from_float_ref +// for this importance-matrix-only codebook type (see +// ../providers/ggml_internal_abi.h's quantize_iq1_s doc comment) -- its +// only quantizer is reached the same way ggml_provider.cpp reaches it: the +// private whole-matrix quantize_iq1_s symbol, with a uniform (all-1.0) +// weighting since it hard-requires a non-null quant_weights. + +#include +#include +#include +#include + +#include + +#include "../providers/ggml_internal_abi.h" +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + ggml_quantize_init(GGML_TYPE_IQ1_S); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_IQ1_S); + const size_t out_bytes = ggml_row_size(GGML_TYPE_IQ1_S, k); + + std::vector weights(k, 1.0f); + std::vector ref_blocks(out_bytes); + quantize_iq1_s(x.data(), ref_blocks.data(), /*nrows=*/1, /*n_per_row=*/k, weights.data()); + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_iq1_s(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_iq2_s.cpp b/apps/ggml/halide/test_iq2_s.cpp new file mode 100644 index 000000000000..ec5bce940667 --- /dev/null +++ b/apps/ggml/halide/test_iq2_s.cpp @@ -0,0 +1,48 @@ +// Standalone round-trip check: the from-scratch Halide IQ2_S dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp). + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + ggml_quantize_init(GGML_TYPE_IQ2_S); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_IQ2_S); + const size_t out_bytes = ggml_row_size(GGML_TYPE_IQ2_S, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_iq2_s(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_iq2_s(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_iq2_xs.cpp b/apps/ggml/halide/test_iq2_xs.cpp new file mode 100644 index 000000000000..05d738095761 --- /dev/null +++ b/apps/ggml/halide/test_iq2_xs.cpp @@ -0,0 +1,47 @@ +// Standalone check: the from-scratch Halide IQ2_XS dequantize kernel vs +// GGML's own reference implementation. GGML has no public from_float_ref +// for this importance-matrix-only codebook type (see +// ../providers/ggml_internal_abi.h's quantize_iq2_xs doc comment) -- its +// only quantizer is reached the same way ggml_provider.cpp reaches it: the +// private whole-matrix quantize_iq2_xs symbol, with a uniform (all-1.0) +// weighting since it hard-requires a non-null quant_weights. + +#include +#include +#include +#include + +#include + +#include "../providers/ggml_internal_abi.h" +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + ggml_quantize_init(GGML_TYPE_IQ2_XS); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_IQ2_XS); + const size_t out_bytes = ggml_row_size(GGML_TYPE_IQ2_XS, k); + + std::vector weights(k, 1.0f); + std::vector ref_blocks(out_bytes); + quantize_iq2_xs(x.data(), ref_blocks.data(), /*nrows=*/1, /*n_per_row=*/k, weights.data()); + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_iq2_xs(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_iq2_xxs.cpp b/apps/ggml/halide/test_iq2_xxs.cpp new file mode 100644 index 000000000000..abc1bfdc9aa7 --- /dev/null +++ b/apps/ggml/halide/test_iq2_xxs.cpp @@ -0,0 +1,52 @@ +// Standalone check: the from-scratch Halide IQ2_XXS dequantize kernel vs +// GGML's own reference implementation. GGML has no public from_float_ref +// for this importance-matrix-only codebook type (see +// ../providers/ggml_internal_abi.h's quantize_iq2_xxs doc comment) -- its +// only quantizer is reached the same way ggml_provider.cpp reaches it: the +// private whole-matrix quantize_iq2_xxs symbol, called with nrows=1 and no +// importance matrix to get a plain per-row reference. + +#include +#include +#include +#include + +#include + +#include "../providers/ggml_internal_abi.h" +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + // The nearest-neighbor grid search this quantizer uses needs its lookup + // table built first (see ggml_provider.cpp's comment on ggml_quantize_init). + ggml_quantize_init(GGML_TYPE_IQ2_XXS); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_IQ2_XXS); + const size_t out_bytes = ggml_row_size(GGML_TYPE_IQ2_XXS, k); + + // quantize_row_iq2_xxs_impl hard-requires a non-null quant_weights + // (GGML_ASSERT) -- a uniform (all-1.0) weighting treats every element + // as equally important, the closest equivalent to "no weighting". + std::vector weights(k, 1.0f); + std::vector ref_blocks(out_bytes); + quantize_iq2_xxs(x.data(), ref_blocks.data(), /*nrows=*/1, /*n_per_row=*/k, weights.data()); + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_iq2_xxs(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_iq3_s.cpp b/apps/ggml/halide/test_iq3_s.cpp new file mode 100644 index 000000000000..2e5132b99486 --- /dev/null +++ b/apps/ggml/halide/test_iq3_s.cpp @@ -0,0 +1,48 @@ +// Standalone round-trip check: the from-scratch Halide IQ3_S dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp). + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + ggml_quantize_init(GGML_TYPE_IQ3_S); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_IQ3_S); + const size_t out_bytes = ggml_row_size(GGML_TYPE_IQ3_S, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_iq3_s(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_iq3_s(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_iq3_xxs.cpp b/apps/ggml/halide/test_iq3_xxs.cpp new file mode 100644 index 000000000000..5dd80636469b --- /dev/null +++ b/apps/ggml/halide/test_iq3_xxs.cpp @@ -0,0 +1,48 @@ +// Standalone round-trip check: the from-scratch Halide IQ3_XXS dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp). + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + ggml_quantize_init(GGML_TYPE_IQ3_XXS); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_IQ3_XXS); + const size_t out_bytes = ggml_row_size(GGML_TYPE_IQ3_XXS, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_iq3_xxs(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_iq3_xxs(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_iq4_nl.cpp b/apps/ggml/halide/test_iq4_nl.cpp new file mode 100644 index 000000000000..456ae55cc725 --- /dev/null +++ b/apps/ggml/halide/test_iq4_nl.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide IQ4_NL dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp). + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK4_NL (32) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_IQ4_NL); + const size_t out_bytes = ggml_row_size(GGML_TYPE_IQ4_NL, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_iq4_nl(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_iq4_nl(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_iq4_xs.cpp b/apps/ggml/halide/test_iq4_xs.cpp new file mode 100644 index 000000000000..628a3a39613a --- /dev/null +++ b/apps/ggml/halide/test_iq4_xs.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide IQ4_XS dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp). + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_IQ4_XS); + const size_t out_bytes = ggml_row_size(GGML_TYPE_IQ4_XS, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_iq4_xs(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_iq4_xs(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_mxfp4.cpp b/apps/ggml/halide/test_mxfp4.cpp new file mode 100644 index 000000000000..9d46c4bfa2df --- /dev/null +++ b/apps/ggml/halide/test_mxfp4.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide MXFP4 dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp). + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_MXFP4 (32) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_MXFP4); + const size_t out_bytes = ggml_row_size(GGML_TYPE_MXFP4, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_mxfp4(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_mxfp4(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_nvfp4.cpp b/apps/ggml/halide/test_nvfp4.cpp new file mode 100644 index 000000000000..3e1e9208dd48 --- /dev/null +++ b/apps/ggml/halide/test_nvfp4.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide NVFP4 dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp). + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_NVFP4 (64) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_NVFP4); + const size_t out_bytes = ggml_row_size(GGML_TYPE_NVFP4, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_nvfp4(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_nvfp4(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q1_0.cpp b/apps/ggml/halide/test_q1_0.cpp new file mode 100644 index 000000000000..a267bbfa5c6e --- /dev/null +++ b/apps/ggml/halide/test_q1_0.cpp @@ -0,0 +1,45 @@ +// Standalone round-trip check: the from-scratch Halide Q1_0 kernels (both +// directions fully native) vs GGML's own reference implementation, reached +// via the public ggml_get_type_traits() API. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK1_0 (128) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q1_0); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q1_0, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q1_0(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q1_0(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q2_k.cpp b/apps/ggml/halide/test_q2_k.cpp new file mode 100644 index 000000000000..0fd2c5deee59 --- /dev/null +++ b/apps/ggml/halide/test_q2_k.cpp @@ -0,0 +1,49 @@ +// Standalone round-trip check: the from-scratch Halide Q2_K dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp), so its +// output is trivially identical to GGML's -- this test's real purpose is +// exercising the from-scratch dequantize implementation against blocks +// GGML itself produced. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q2_K); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q2_K, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q2_k(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q2_k(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q3_k.cpp b/apps/ggml/halide/test_q3_k.cpp new file mode 100644 index 000000000000..767eb193727b --- /dev/null +++ b/apps/ggml/halide/test_q3_k.cpp @@ -0,0 +1,49 @@ +// Standalone round-trip check: the from-scratch Halide Q3_K dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp), so its +// output is trivially identical to GGML's -- this test's real purpose is +// exercising the from-scratch dequantize implementation against blocks +// GGML itself produced. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q3_K); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q3_K, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q3_k(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q3_k(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q4_0.cpp b/apps/ggml/halide/test_q4_0.cpp new file mode 100644 index 000000000000..f568b0f17401 --- /dev/null +++ b/apps/ggml/halide/test_q4_0.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide Q4_0 kernels vs +// GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API (same accessor providers/ggml_provider.cpp +// uses). Run before wiring this into kernel-bench as a provider. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK4_0 (32) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q4_0); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q4_0, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q4_0(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q4_0(halide_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q4_1.cpp b/apps/ggml/halide/test_q4_1.cpp new file mode 100644 index 000000000000..bdbb4f9f0c5d --- /dev/null +++ b/apps/ggml/halide/test_q4_1.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide Q4_1 kernels vs +// GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API (same accessor providers/ggml_provider.cpp +// uses). Run before wiring this into kernel-bench as a provider. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK4_1 (32) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q4_1); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q4_1, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q4_1(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q4_1(halide_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q4_k.cpp b/apps/ggml/halide/test_q4_k.cpp new file mode 100644 index 000000000000..62174d0e5e3d --- /dev/null +++ b/apps/ggml/halide/test_q4_k.cpp @@ -0,0 +1,49 @@ +// Standalone round-trip check: the from-scratch Halide Q4_K dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp), so its +// output is trivially identical to GGML's -- this test's real purpose is +// exercising the from-scratch dequantize implementation against blocks +// GGML itself produced. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q4_K); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q4_K, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q4_k(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q4_k(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q5_0.cpp b/apps/ggml/halide/test_q5_0.cpp new file mode 100644 index 000000000000..d64bf692ff80 --- /dev/null +++ b/apps/ggml/halide/test_q5_0.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide Q5_0 kernels vs +// GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API (same accessor providers/ggml_provider.cpp +// uses). Run before wiring this into kernel-bench as a provider. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK5_0 (32) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q5_0); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q5_0, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q5_0(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q5_0(halide_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q5_1.cpp b/apps/ggml/halide/test_q5_1.cpp new file mode 100644 index 000000000000..2b096dfa3ebd --- /dev/null +++ b/apps/ggml/halide/test_q5_1.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide Q5_1 kernels vs +// GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API (same accessor providers/ggml_provider.cpp +// uses). Run before wiring this into kernel-bench as a provider. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK5_1 (32) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q5_1); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q5_1, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q5_1(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q5_1(halide_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q5_k.cpp b/apps/ggml/halide/test_q5_k.cpp new file mode 100644 index 000000000000..d54e62c4f9dc --- /dev/null +++ b/apps/ggml/halide/test_q5_k.cpp @@ -0,0 +1,49 @@ +// Standalone round-trip check: the from-scratch Halide Q5_K dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp), so its +// output is trivially identical to GGML's -- this test's real purpose is +// exercising the from-scratch dequantize implementation against blocks +// GGML itself produced. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q5_K); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q5_K, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q5_k(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q5_k(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q6_k.cpp b/apps/ggml/halide/test_q6_k.cpp new file mode 100644 index 000000000000..804fb2f35081 --- /dev/null +++ b/apps/ggml/halide/test_q6_k.cpp @@ -0,0 +1,49 @@ +// Standalone round-trip check: the from-scratch Halide Q6_K dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp), so its +// output is trivially identical to GGML's -- this test's real purpose is +// exercising the from-scratch dequantize implementation against blocks +// GGML itself produced. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q6_K); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q6_K, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q6_k(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q6_k(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q8_0.cpp b/apps/ggml/halide/test_q8_0.cpp new file mode 100644 index 000000000000..95135acf835b --- /dev/null +++ b/apps/ggml/halide/test_q8_0.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide Q8_0 kernels vs +// GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API (same accessor providers/ggml_provider.cpp +// uses). Run before wiring this into kernel-bench as a provider. + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK8_0 (32) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q8_0); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q8_0, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q8_0(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_q8_0(halide_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q8_1.cpp b/apps/ggml/halide/test_q8_1.cpp new file mode 100644 index 000000000000..fe570667642f --- /dev/null +++ b/apps/ggml/halide/test_q8_1.cpp @@ -0,0 +1,37 @@ +// Standalone check: the from-scratch Halide Q8_1 quantize kernel vs GGML's +// own reference implementation, reached via the public +// ggml_get_type_traits() API. Q8_1 is an activation-only format -- GGML has +// no public to_float for it, so there's no dequantize round-trip to check +// here, only the quantize output. + +#include +#include +#include +#include + +#include + +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK8_1 (32) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_Q8_1); + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q8_1, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q8_1(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_q8_k.cpp b/apps/ggml/halide/test_q8_k.cpp new file mode 100644 index 000000000000..5413714382c5 --- /dev/null +++ b/apps/ggml/halide/test_q8_k.cpp @@ -0,0 +1,40 @@ +// Standalone check: the from-scratch Halide Q8_K quantize kernel vs GGML's +// own reference implementation. Unlike every other type here, Q8_K has no +// public from_float_ref (see include/ggml.h's type_traits table) -- it's +// reached the same way apps/ggml/providers/ggml_provider.cpp reaches it, +// through the private quantize_row_q8_K_generic symbol declared in +// ../providers/ggml_internal_abi.h (see that header's own comment for why +// this one symbol needs the private ABI). Q8_K is activation-only, so +// there's no dequantize round-trip to check, only the quantize output. + +#include +#include +#include +#include + +#include + +#include "../providers/ggml_internal_abi.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const size_t out_bytes = ggml_row_size(GGML_TYPE_Q8_K, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + quantize_row_q8_K_generic(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_q8_k(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_tq1_0.cpp b/apps/ggml/halide/test_tq1_0.cpp new file mode 100644 index 000000000000..3c08bb576797 --- /dev/null +++ b/apps/ggml/halide/test_tq1_0.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide TQ1_0 dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp). + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_TQ1_0); + const size_t out_bytes = ggml_row_size(GGML_TYPE_TQ1_0, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_tq1_0(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_tq1_0(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/test_tq2_0.cpp b/apps/ggml/halide/test_tq2_0.cpp new file mode 100644 index 000000000000..19464cb6eab6 --- /dev/null +++ b/apps/ggml/halide/test_tq2_0.cpp @@ -0,0 +1,46 @@ +// Standalone round-trip check: the from-scratch Halide TQ2_0 dequantize +// kernel vs GGML's own reference implementation, reached via the public +// ggml_get_type_traits() API. Quantize here is scaffolding that itself +// calls out to GGML's reference (see ggml_extern_quantize.cpp). + +#include +#include +#include +#include + +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_quants.h" + +int main() { + const int64_t k = 4096; // multiple of QK_K (256) + + std::vector x(k); + generate_synthetic_data(x.data(), k); + + const ggml_type_traits *tt = ggml_get_type_traits(GGML_TYPE_TQ2_0); + const size_t out_bytes = ggml_row_size(GGML_TYPE_TQ2_0, k); + + std::vector ref_blocks(out_bytes), halide_blocks(out_bytes); + tt->from_float_ref(x.data(), ref_blocks.data(), k); + ggml_quants_halide_quantize_tq2_0(x.data(), halide_blocks.data(), k); + + if (std::memcmp(ref_blocks.data(), halide_blocks.data(), out_bytes) != 0) { + std::fprintf(stderr, "FAIL: quantize output does not match GGML's reference byte-for-byte\n"); + return 1; + } + + std::vector ref_y(k), halide_y(k); + tt->to_float(ref_blocks.data(), ref_y.data(), k); + ggml_quants_halide_dequantize_tq2_0(ref_blocks.data(), halide_y.data(), k); + + if (!floats_match(ref_y.data(), halide_y.data(), k)) { + std::fprintf(stderr, "FAIL: dequantize output does not match GGML's reference within tolerance\n"); + return 1; + } + + std::printf("Success!\n"); + return 0; +} diff --git a/apps/ggml/halide/vec_dot_generator_base.h b/apps/ggml/halide/vec_dot_generator_base.h new file mode 100644 index 000000000000..e95ff49a27c9 --- /dev/null +++ b/apps/ggml/halide/vec_dot_generator_base.h @@ -0,0 +1,584 @@ +#pragma once + +// Shared configure() scaffolding for every Approximation-based vec_dot +// Generator (the extended SymmetricVecDotGenerator, KQuantVecDotGenerator, +// LookupTableVecDotGenerator). All three build the same "naive fp32 dot +// product -> approximate_by both operands -> sever severs the +// (already-quantized, Input-supplied) encode halves -> schedule" pipeline; +// they differ only in which block-indexed codecs their build_vec_dot() picks. +// This factors that shared body out via CRTP (Derived::build_vec_dot()) -- the +// same static-polymorphism idiom as codec_generator_base.h's +// CodecGeneratorBase. +// +// generate() never calls Approximation::encode()/decode() directly -- only +// through Func::approximate_by()/Pipeline::sever(), exactly like the +// codec generators. The vec_dot is the point at which the framework's splice + +// sever path is exercised for a dot product (matching +// test/performance/matvec_offline_split.cpp). +// +// Usage: +// class FooVecDotGenerator : public VecDotGeneratorBase { +// public: +// GeneratorParam<...> family{...}; +// VecDotSpec build_vec_dot() const { return {weight_codec, wbytes, act_codec, abytes, block_size, sched}; } +// }; + +#include "Halide.h" + +#include "codec_generator_base.h" // Direction/SchemeAndBytes live here; shared idiom +#include "quant_components.h" +#include "sdot_schedule.h" + +namespace ggml_halide { + +// Whether the per-block scale factors out to a single block-invariant scalar +// (SDOT: hoist_invariants() + rfactor() + change_type(Int(32)) -> Int(32) inner +// dot) or not (Float: a plain vectorized float accumulation -- affine offsets, +// two-level sub-block scales, and per-group grid scales are not +// single-per-block-invariant). +enum class ScheduleKind { SDOT, + Float }; + +struct VecDotSpec { + // Both codecs decode to a block-indexed (kk, blk) Func at the SAME + // block_size, so the reduction below is uniform -- Wt(r.x, r.y) * Vec(r.x, + // r.y). When an activation is stored in a smaller block than the weight + // (e.g. Q1_0/NVFP4 x Q8_0), its codec is composed with a Reblock stage + // (see quant_components.h) that re-views it at the weight's block_size -- + // the block-structure reconciliation is an Approximation, not something + // the Generator open-codes into the reduction. + Halide::Approximation weight_codec; + int weight_bytes; + Halide::Approximation act_codec; + int act_bytes; + int block_size; + ScheduleKind sched; + // Whether the SDOT schedule must multiply the per-block product out before + // hoisting. Set for the formats whose weight decode carries an offset (the + // affine family): (d*code + m) * (d_act*act) has no single scale to hoist + // until it is expanded. See sdot_schedule.h. + bool distribute_terms; + // When set (is_struct()), the corresponding operand's packed blocks are a + // first-class 1-D Type::Struct buffer (one struct per block) rather than a + // 2-D (byte, blk) UInt(8) buffer. Default-invalid keeps a codec on the byte + // path -- so an operand whose codec isn't struct-typed (e.g. a Reblock'd + // activation) just leaves its type unset. + Halide::Type weight_type; + Halide::Type act_type; + // Set when the activation stores a per-block scaled sum of its codes (Q8_1's + // `s` field). Together with distribute_terms (an affine weight), this lets + // configure() sever the offset term's sum(act) accumulator to a stored fp16 + // field supplied as a third Input, instead of recomputing it -- ggml's own + // q4_1/q5_1 optimization. See SchemeAndBytes::has_block_sums. + bool act_has_block_sums = false; + // Number of adjacent blocks kept in flight by the SDOT schedule. The + // ordinary q4/q8 paths prefer four; q5's table-expansion live ranges make + // two faster and match GGML's paired-block loop. + int unroll_blocks = 4; + Halide::Approximation reconstructed_codes_stage; + Halide::Approximation packed_high_word_stage; + // Clean core-composed schemes can share one traced decode graph between + // the four-block main update and its tiny remainder update. Legacy schemes + // retain separate traces until their representation stages are migrated. + bool share_weight_tail = false; + bool share_act_tail = false; +}; + +template +class VecDotGeneratorBase : public Halide::Generator { +public: + void configure() { + using namespace Halide; + VecDotSpec spec = static_cast(this)->build_vec_dot(); + int bs = spec.block_size; + const int unroll_blocks = spec.unroll_blocks; + + // A struct-typed operand's packed blocks are a 1-D Type::Struct buffer + // (block index only); a byte-path operand is 2-D (byte, blk). + const bool wt_struct = spec.weight_type.is_struct(); + const bool act_struct = spec.act_type.is_struct(); + ImageParam x_blocks = wt_struct ? ImageParam(spec.weight_type, 1, "x_blocks") : ImageParam(UInt(8), 2, "x_blocks"); + ImageParam y_blocks = act_struct ? ImageParam(spec.act_type, 1, "y_blocks") : ImageParam(UInt(8), 2, "y_blocks"); + + // Third Input, present only for the affine-x-Q8_1 pairings: the + // activation's stored per-block scaled code sum (`s`). configure() severs + // the offset term's sum(act) accumulator to this field rather than + // recomputing it -- a zero-copy 1-D fp16 view of the `s` slot the ABI + // wrapper passes (stride = block width). See the SDOT sever branch below. + const bool sever_sum = spec.sched == ScheduleKind::SDOT && + spec.distribute_terms && spec.act_has_block_sums; + ImageParam s_blocks = sever_sum ? ImageParam(Float(16), 1, "s_blocks") : ImageParam(); + + // The packed-block buffers are quantized GGML rows: their base pointers + // are cache-line aligned. Without this Halide assumes 1-byte alignment + // and lowers every strided / reinterpreted read (the interleaved fp16 + // scales, the int8 codes) to byte-wise ld1.b + orr reassembly instead + // of wide vector loads -- which dominates these tiny vec_dots. + x_blocks.set_host_alignment(64); + y_blocks.set_host_alignment(64); + + // Naive fp32 placeholders -- never realized; sever() severs + // Acc from them entirely, and the real values come from the + // already-quantized x_blocks/y_blocks. Block-indexed (kk, blk) to match + // the codecs' block-indexed decode. + Var kk("kk"), blk("blk"), u("u"); + Func Wt("wt_naive"), Vec("vec_naive"); + Wt(kk, blk) = 0.0f; + Vec(kk, blk) = 0.0f; + + // The SDOT schedule interleaves unroll_blocks blocks into independent + // accumulators, so it wants a block count divisible by that. Letting + // Halide's split produce the odd tail instead is not a local cost: the + // predicate it inserts makes the per-block sdot a dynamic-extent + // allocation that has to be zeroed and accumulated through memory, + // roughly doubling the cost of *every* block. So the main reduction gets + // an exactly divisible extent and a second update sweeps the remainder + // at the default schedule (at most unroll_blocks - 1 blocks). + const bool sdot = spec.sched == ScheduleKind::SDOT; + const bool keyed_q5 = spec.reconstructed_codes_stage.defined(); + const bool share_weight_tail = keyed_q5 || spec.share_weight_tail; + const bool share_act_tail = spec.share_act_tail; + Expr nblocks = x_blocks.dim(wt_struct ? 0 : 1).extent(); + Expr main_blocks = sdot ? (nblocks / unroll_blocks) * unroll_blocks : nblocks; + + RDom r(0, bs, 0, main_blocks, "r"); + Func Acc("acc"); + Acc() = 0.0f; + Acc() += Wt(r.x, r.y) * Vec(r.x, r.y); + + // The odd-block tail decodes through its OWN placeholder Funcs, given a + // separate decode chain by a second approximate_by below. This lets the + // main reduction materialize a reconstructed-codes leaf (Q5_x) via + // compute_at while the tail -- a different loop nest that could not see + // that per-block buffer -- reconstructs inline, at negligible cost (fewer + // than unroll_blocks blocks). Both chains read the same x_blocks/y_blocks. + Func WtT("wt_naive_tail"), VecT("vec_naive_tail"); + WtT(kk, blk) = 0.0f; + VecT(kk, blk) = 0.0f; + RDom r_tail(0, bs, main_blocks, nblocks - main_blocks, "r_tail"); + if (sdot) { + // q5_0 shares the main decode graph with its tiny odd-block tail. + // The tail update is eagerly inlined below before the reconstructed + // codes leaf is materialized for the main paired-block update. + Func tail_weight = share_weight_tail ? Wt : WtT; + Func tail_act = share_act_tail ? Vec : VecT; + Acc() += tail_weight(r_tail.x, r_tail.y) * tail_act(r_tail.x, r_tail.y); + } + + ApproximationResult wt_r = Wt.approximate_by(spec.weight_codec, {Acc}); + ApproximationResult act_r = Vec.approximate_by(spec.act_codec, {Acc}); + + // Both operands' encode halves are severed and bound to the real + // already-quantized Input buffers (same as symmetric_vec_dot). For the + // extern-delegated / SeveredEncode weight schemes the encode is likewise + // severed here, so its extern symbol is never computed or linked. + std::vector to_sever = wt_r.encoded; + to_sever.insert(to_sever.end(), act_r.encoded.begin(), act_r.encoded.end()); + std::vector bind_to = {x_blocks, y_blocks}; + ApproximationResult wtT_r, actT_r; + if (sdot) { + if (!share_weight_tail) { + wtT_r = WtT.approximate_by(spec.weight_codec, {Acc}); + to_sever.insert(to_sever.end(), wtT_r.encoded.begin(), wtT_r.encoded.end()); + bind_to.push_back(x_blocks); + for (Func h : wtT_r.intermediates) { + if (h.has_update_definition()) { + h.compute_root(); + } + } + } + if (!share_act_tail) { + actT_r = VecT.approximate_by(spec.act_codec, {Acc}); + to_sever.insert(to_sever.end(), actT_r.encoded.begin(), actT_r.encoded.end()); + bind_to.push_back(y_blocks); + for (Func h : actT_r.intermediates) { + if (h.has_update_definition()) { + h.compute_root(); + } + } + } + } + Pipeline({Acc}).sever(to_sever, bind_to); + + // Q5_0/Q5_1 reconstruct each code from a nibble plus a per-element high + // bit read from the qh field's byte->bits expansion table (see + // PlanarBitPack::decode). That table read is only a *contiguous* 8-byte + // load -- matching ggml's table_b2b -- when the qh byte is a scalar and + // the 8 bit positions are the vector lanes. Inlined into the sdot it is + // the opposite (qh byte per lane -> a per-lane gather), so the SDOT + // branches materialize the reconstructed int8 codes per block (compute_at + // the block loop, kk split as (byte, pos): pos vectorizes the table load, + // byte unrolls to a scalar index). The odd-block tail decodes through its + // own inline chain (see above), so it does not need this buffer. + Func codes_leaf, qh_leaf; + if (keyed_q5) { + codes_leaf = wt_r.decoded_by(spec.reconstructed_codes_stage); + qh_leaf = wt_r.decoded_by(spec.packed_high_word_stage); + _halide_internal_assert(codes_leaf.defined() && qh_leaf.defined()); + } else { + // q5_1's legacy composition has no stage handles for these Funcs + // (qh is an intermediate of Q5StructBlockLayout, not a port), + // so its boundaries are still discovered by generated name. + for (const Func &h : wt_r.intermediates) { + if (h.name() == "combine_bits_code") { + codes_leaf = h; + } else if (h.name() == "q5_struct_block_qh") { + qh_leaf = h; + } + } + } + + if (share_weight_tail || share_act_tail) { + // Flatten shared decode stages into this update only. The main + // update is scheduled independently below, so its SDOT boundaries + // and any keyed materializations remain intact. + std::vector tail_inline; + if (share_weight_tail) { + tail_inline.push_back(wt_r.replacement); + for (const Func &h : wt_r.intermediates) { + if (h.function().can_be_inlined()) { + tail_inline.push_back(h); + } + } + } + if (share_act_tail) { + tail_inline.push_back(act_r.replacement); + for (const Func &h : act_r.intermediates) { + if (h.function().can_be_inlined()) { + tail_inline.push_back(h); + } + } + } + for (size_t pass = 0; pass < tail_inline.size(); ++pass) { + Acc.update(1).eager_inline(tail_inline); + } + } + auto schedule_codes = [&](LoopLevel level) { + // Split kk into (byte, pos): pos (8) vectorizes the contiguous table + // load, byte unrolls to a scalar index. store_in(Register) keeps the + // reconstructed codes in the vector register file straight into the + // sdot instead of round-tripping a stack buffer (all stores are at + // constant coordinates once ki is vectorized and ko unrolled). + Var kc = codes_leaf.args()[0], co("co"), ci("ci"), byte("byte"), pos("pos"); + codes_leaf.compute_at(level) + .store_in(MemoryType::Register) + .split(kc, co, ci, 16) // 16-code register unit = one sdot chunk + .split(ci, byte, pos, 8) // within it, two qh bytes x 8 positions + .vectorize(pos, 8) + .unroll(byte) + .unroll(co); + if (qh_leaf.defined()) { + qh_leaf.compute_at(level).store_in(MemoryType::Register); + } + }; + + // Only intermediates with update definitions (per-block stat reductions) need + // explicit scheduling; pure pass-throughs stay inline (same reasoning as + // symmetric_vec_dot_generator.cpp). + for (Func h : wt_r.intermediates) { + if (h.has_update_definition()) { + h.compute_root(); + } + } + for (Func h : act_r.intermediates) { + if (h.has_update_definition()) { + h.compute_root(); + } + } + + Func final_value = Acc; + if (sever_sum && codes_leaf.defined()) { + // Q5_1 needs two differently shaped reductions: a four-lane dot + // accumulator that survives across blocks, and the scalar m*s + // term supplied by Q8_1's stored sum. Express them independently, + // then fuse their paired-block loops. This mirrors GGML and avoids + // the generic sever path's horizontal Int(32) reduction per block. + Func min_leaf, scale_leaf; + for (const Func &h : wt_r.intermediates) { + if (h.name() == "q5_struct_block_min") { + min_leaf = h; + } else if (h.name() == "q5_struct_block_scale") { + scale_leaf = h; + } + } + Func codes_tail, scale_tail; + for (const Func &h : wtT_r.intermediates) { + if (h.name().find("combine_bits_code") == 0) { + codes_tail = h; + } else if (h.name().find("q5_struct_block_scale") == 0) { + scale_tail = h; + } + } + _halide_internal_assert(min_leaf.defined() && scale_leaf.defined() && + codes_tail.defined() && scale_tail.defined()); + + Func wt_product("q5_1_product_weight"); + wt_product(kk, blk) = cast(codes_leaf(kk, blk)) * scale_leaf(blk); + Func wt_product_tail("q5_1_product_weight_tail"); + wt_product_tail(kk, blk) = cast(codes_tail(kk, blk)) * scale_tail(blk); + + RDom rd(0, bs, 0, main_blocks, "rd"); + RDom rd_tail(0, bs, main_blocks, nblocks - main_blocks, "rd_tail"); + Func dot_acc("q5_1_dot_acc"); + dot_acc() = 0.0f; + dot_acc() += wt_product(rd.x, rd.y) * act_r.replacement(rd.x, rd.y); + dot_acc() += wt_product_tail(rd_tail.x, rd_tail.y) * actT_r.replacement(rd_tail.x, rd_tail.y); + + const int lanes = 4; + RVar rxc("rxc"), rxr("rxr"), rxo("rxo"), rxi("rxi"); + dot_acc.update(0).split(rd.x, rxc, rxr, 4 * lanes); + dot_acc.update(0).split(rxr, rxo, rxi, 4); + dot_acc.update(0).eager_inline(wt_product); + Var lane("lane"); + std::vector dot_i32 = sdot_partial(dot_acc, {{rxo, lane}, {rd.y, u}}, + {act_r}, false, {codes_leaf}); + + RVar ryo("ryo"), ryi("ryi"); + Var lv("lv"), bacc("bacc"); + dot_acc.update(0).split(rd.y, ryo, ryi, unroll_blocks); + Func acc_vec = dot_acc.update(0).rfactor({{rxo, lv}, {ryi, bacc}}); + acc_vec.compute_root().vectorize(lv, lanes).unroll(bacc); + acc_vec.update().vectorize(lv, lanes).unroll(bacc); + for (Func &part : dot_i32) { + part.compute_at(acc_vec, bacc) + .update() + .atomic() + .vectorize(rxi, 4) + .vectorize(lane, lanes) + .unroll(rxc); + } + schedule_codes(LoopLevel(acc_vec, bacc)); + + Var lv2("lv2"); + Func dot_lanes = dot_acc.update(0).rfactor(rxo, lv2); + dot_lanes.compute_root().vectorize(lv2, lanes); + dot_lanes.update().vectorize(lv2, lanes).unroll(ryi); + dot_acc.update(0).atomic().vectorize(rxo, lanes); + dot_acc.update(1).unscheduled(); + + RDom rb(0, main_blocks, "rb"); + RDom rb_tail(main_blocks, nblocks - main_blocks, "rb_tail"); + Func offset_acc("q5_1_offset_acc"); + offset_acc() = 0.0f; + offset_acc() += min_leaf(rb) * cast(s_blocks(rb)); + offset_acc() += min_leaf(rb_tail) * cast(s_blocks(rb_tail)); + + RVar rbi("rbi"); + Var offset_bacc("offset_bacc"); + offset_acc.update(0).split(rb.x, ryo, rbi, unroll_blocks); + Func offset_vec = offset_acc.update(0).rfactor(rbi, offset_bacc); + offset_vec.compute_root().unroll(offset_bacc); + offset_vec.update().unroll(offset_bacc); + offset_vec.update().compute_with(acc_vec.update(), ryo, LoopAlignStrategy::AlignStart); + offset_acc.update(0).unroll(rbi); + offset_acc.update(1).unscheduled(); + + Func combined("q5_1_combined_acc"); + combined() = dot_acc() + offset_acc(); + final_value = combined; + } else if (sever_sum) { + // Affine weight x Q8_1: the per-block product (d*code + m)*(d_act*act) + // distributes into d*d_act*sum(code*act) + m*d_act*sum(act). ggml does + // not recompute the second sum -- it reads the `s` field Q8_1 stores. + // We do the same by severing that term's accumulator to the stored + // field (the third Input, s_blocks), leaving only the Int(32) dot. + // + // rfactor to whole-block partials (variant A), inlining only the + // WEIGHT's decode chain so the activation decode (Act) stays whole: + // the offset term's accumulator is then sum_k Act(k, blk), which *is* + // the stored `s`. The product term re-inlines Act's full chain and + // re-hoists to recover the scale-free Int(32) dot. + Func acc_dot = Acc.update().rfactor({{r.y, u}}); + + std::vector winl = {wt_r.replacement}; + for (const Func &h : wt_r.intermediates) { + // Keep the materialized codes leaf (Q5_1) out of the flatten, so + // its qh table read stays a per-block contiguous load. + if (h.function().can_be_inlined() && + !(codes_leaf.defined() && h.name() == codes_leaf.name())) { + winl.push_back(h); + } + } + for (size_t pass = 0; pass < winl.size(); pass++) { + acc_dot.update().eager_inline(winl); + } + + acc_dot.update().distribute(); + std::vector parts = acc_dot.update().hoist_invariants(); + // parts[0] = product term (scale * sum code_w*Act); parts[1] = offset + // term (min * sum Act == stored s). + + // Sever the offset term to the stored fp16 field. change_type(Float16) + // makes the severed accumulator's type match the data (the encoder + // rounds `s` to fp16 too, so this reproduces ggml's own rounding); + // sever then replaces every call to it with a read of + // s_blocks and discards the recomputing reduction. + Func s16 = parts[1].change_type(Float(16)); + Pipeline({Acc}).sever({s16}, {s_blocks}); + + // Product term: flatten the activation's full decode chain and + // re-hoist to pull d_act out, leaving the scale-free Int(32) dot. + std::vector ainl = {act_r.replacement}; + for (const Func &h : act_r.intermediates) { + if (h.function().can_be_inlined()) { + ainl.push_back(h); + } + } + for (size_t pass = 0; pass < ainl.size(); pass++) { + parts[0].update().eager_inline(ainl); + } + Func prod_i32 = parts[0].update().hoist_invariants()[0].change_type(Int(32)); + + RVar ryo("ryo"), ryi("ryi"); + Var bacc("bacc"); + Acc.update(0).split(r.y, ryo, ryi, unroll_blocks); + Func acc_vec = Acc.update(0).rfactor(ryi, bacc); + acc_vec.compute_root().unroll(bacc); + acc_vec.update().unroll(bacc); + + // Only the product dot is computed per block now; the offset term is + // a severed read of s_blocks (nothing to schedule). + prod_i32.compute_at(acc_vec, bacc).update().atomic().vectorize(r.x, bs); + if (codes_leaf.defined()) { + schedule_codes(LoopLevel(acc_vec, bacc)); + } + Acc.update(1).unscheduled(); + } else if (spec.sched == ScheduleKind::SDOT && getenv("GGML_PER_BLOCK_PROBE")) { + // PROBE (variant A): rfactor only the block index, so both terms are + // whole-block reductions. Costs a horizontal reduce per block and a + // scalar cross-block accumulator; kept for measuring the non-sum + // formats against the lane-split default below. + std::vector parts = sdot_partial(Acc, {{r.y, u}}, {wt_r, act_r}, spec.distribute_terms); + + RVar ryo("ryo"), ryi("ryi"); + Var bacc("bacc"); + Acc.update(0).split(r.y, ryo, ryi, unroll_blocks); + Func acc_vec = Acc.update(0).rfactor(ryi, bacc); + acc_vec.compute_root().unroll(bacc); + acc_vec.update().unroll(bacc); + + for (Func &part : parts) { + part.compute_at(acc_vec, bacc).update().atomic().vectorize(r.x, bs); + } + Acc.update(1).unscheduled(); + } else if (spec.sched == ScheduleKind::SDOT) { + // The reduction is over (within-block r.x) x (block r.y). The lanes + // of the accumulator come from r.x, so the sdot's four Int(32) lanes + // survive all the way into the float accumulator and no block pays + // for a horizontal reduce. They must come from r.x rather than r.y: + // blocks are interleaved {scale, codes} records, so a lane per block + // would gather both the codes and the scales, while a lane per + // within-block group keeps every code load contiguous. + // + // One sdot consumes 16 int8s and lands in a 4-lane Int(32) register, + // so r.x is cut three ways: chunks of 16 (one sdot each, run + // serially so they accumulate into the *same* register), then within + // a chunk a 4-wide lane index and the 4 elements the lane sums. + // Reducing straight to 4 lanes instead would make Halide lower the + // wide reduce as two independent sdots plus an addp to merge them -- + // an extra zeroing and an extra reduction per block. + const int lanes = 4; + const int chunk = 4 * lanes; + RVar rxc("rxc"), rxr("rxr"), rxo("rxo"), rxi("rxi"); + Acc.update(0).split(r.x, rxc, rxr, chunk); + Acc.update(0).split(rxr, rxo, rxi, 4); + + // sdot_partial() flattens both operands' decode chains and hoists + // their per-block scales out of the surviving rxi reduction, leaving + // the scale-free Int(32) dot. See sdot_schedule.h. + Var lane("lane"); + std::vector keep_out; + if (codes_leaf.defined()) { + keep_out.push_back(codes_leaf); + } + std::vector Acc_i32 = sdot_partial(Acc, {{rxo, lane}, {r.y, u}}, {wt_r, act_r}, spec.distribute_terms, keep_out); + + // Acc's update now reduces over (rxo, r.y). Peel rxo back off as the + // vector lanes, and peel unroll_blocks consecutive blocks off + // alongside it into separate accumulators. The accumulators have to + // be split over *blocks*: every lane of one accumulator advances on + // every block, so widening the vector does not shorten the + // multiply-add chain, only interleaving blocks does. At ~3-4 cycles + // of accumulate latency, an un-interleaved chain is what bounds the + // whole kernel. + RVar ryo("ryo"), ryi("ryi"); + Var lv("lv"), bacc("bacc"); + Acc.update(0).split(r.y, ryo, ryi, unroll_blocks); + Func acc_vec = Acc.update(0).rfactor({{rxo, lv}, {ryi, bacc}}); + acc_vec.compute_root().vectorize(lv, lanes).unroll(bacc); + acc_vec.update().vectorize(lv, lanes).unroll(bacc); + + // Inside the unrolled body, not at the block-group loop: at `bacc` the + // sdot is one block's worth of registers, whereas at `ryo` it is a + // unroll_blocks-long buffer that Halide has to allocate, zero, and + // accumulate through memory. + for (Func &part : Acc_i32) { + part.compute_at(acc_vec, bacc) + .update() + .atomic() + .vectorize(rxi, 4) + .vectorize(lane, lanes) + .unroll(rxc); + } + if (codes_leaf.defined()) { + schedule_codes(LoopLevel(acc_vec, bacc)); + } + + // Collapsing the lanes x unrolled-blocks accumulators is a fixed + // cost, but at the row lengths GGML uses it is not a negligible one: + // left alone it is a serial chain of lanes*unroll_blocks scalar + // adds. Sum the blocks vectorially first, then reduce the lanes + // horizontally, so it costs a handful of vector ops instead. + Var lv2("lv2"); + Func acc_lanes = Acc.update(0).rfactor(rxo, lv2); + acc_lanes.compute_root().vectorize(lv2, lanes); + acc_lanes.update().vectorize(lv2, lanes).unroll(ryi); + Acc.update(0).atomic().vectorize(rxo, lanes); + + // The odd-block tail deliberately keeps the default schedule. + Acc.update(1).unscheduled(); + } + // ScheduleKind::Float: leave the reduction at its default (legal) schedule + // -- correctness first; an interleave/sub-block-aware performance schedule + // is a separate step. + + Func result("result"); + result() = final_value(); + + // A byte-path operand's block stride is pinned to its block width: these + // are densely packed GGML rows, and leaving the stride dynamic costs a + // serial pointer-add chain per block instead of an immediate offset. + if (wt_struct) { + x_blocks.dim(0).set_min(0); + } else { + x_blocks.dim(0).set_bounds(0, spec.weight_bytes); + x_blocks.dim(1).set_min(0).set_stride(spec.weight_bytes); + } + if (act_struct) { + y_blocks.dim(0).set_min(0); + } else { + y_blocks.dim(0).set_bounds(0, spec.act_bytes); + y_blocks.dim(1).set_min(0).set_stride(spec.act_bytes); + } + if (sever_sum) { + // A gathered view of the `s` slot within each packed block: the field + // repeats every act_bytes, so the fp16 stride is act_bytes/2 (18 for + // Q8_1's 36-byte block). Pinning it makes the per-block read a + // compile-time immediate offset rather than a dynamic pointer chain. + s_blocks.dim(0).set_min(0).set_stride(spec.act_bytes / 2); + } + + this->add_input(x_blocks); + this->add_input(y_blocks); + if (sever_sum) { + this->add_input(s_blocks); + } + this->add_output(result); + } + + void generate() { + // configure() built the whole pipeline (add_input/add_output included). + } +}; + +} // namespace ggml_halide diff --git a/apps/ggml/include/kernel_registry.h b/apps/ggml/include/kernel_registry.h new file mode 100644 index 000000000000..34ce7d3acb1b --- /dev/null +++ b/apps/ggml/include/kernel_registry.h @@ -0,0 +1,156 @@ +#pragma once + +// Implementation-agnostic core of kernel-bench. +// +// A "kernel" here is identified by a ggml_type plus a category (quantize, +// dequantize, vec_dot, repack quantize_mat/gemv/gemm). For each (category, +// type) pair, exactly one *reference* implementation may be registered -- +// this is the correctness ground truth and baseline timing that every other +// *candidate* implementation for that (category, type) is measured against. +// +// GGML's own routines are registered as the reference by providers/ggml_provider.cpp. +// Nothing in this file knows anything about GGML internals: a future provider +// (e.g. the user's own reference implementation) is just another call to +// register_candidate() (or register_reference(), if it should replace GGML's +// as the ground truth) with a function pointer matching the category's +// signature. See providers/README.md for how to add one. + +#include +#include +#include +#include +#include +#include + +#include + +// Function-pointer shapes shared by every provider. These mirror the layouts +// used throughout ggml-cpu (quantize_row_*, ggml_vec_dot_*, ggml_gemv_*/ggml_gemm_*) +// so that both GGML's own symbols and a from-scratch implementation can be +// registered without adapters. +using quantize_fn_t = void (*)(const float *GGML_RESTRICT x, void *GGML_RESTRICT y, int64_t k); +using dequantize_fn_t = void (*)(const void *GGML_RESTRICT x, float *GGML_RESTRICT y, int64_t k); +using vec_dot_fn_t = void (*)(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); +using gemx_fn_t = void (*)(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, + const void *GGML_RESTRICT vy, int nr, int nc); + +template +struct Impl { + std::string name; + Fn fn; +}; + +template +class Registry { +public: + // Exactly one reference per type. Calling this twice for the same type + // replaces the previous reference (last call wins) -- useful if a later + // provider should become the new ground truth for that type. + void register_reference(ggml_type type, std::string name, Fn fn) { + entries_[type].reference = Impl{std::move(name), fn}; + } + + // Zero or more per type. + void register_candidate(ggml_type type, std::string name, Fn fn) { + entries_[type].candidates.push_back(Impl{std::move(name), fn}); + } + + const Impl *reference(ggml_type type) const { + auto it = entries_.find(type); + if (it == entries_.end() || !it->second.reference.has_value()) { + return nullptr; + } + return &*it->second.reference; + } + + const std::vector> &candidates(ggml_type type) const { + static const std::vector> empty; + auto it = entries_.find(type); + return it == entries_.end() ? empty : it->second.candidates; + } + + std::vector types_with_reference() const { + std::vector out; + for (const auto &[type, entry] : entries_) { + if (entry.reference.has_value()) { + out.push_back(type); + } + } + return out; + } + +private: + struct Entry { + std::optional> reference; + std::vector> candidates; + }; + std::map entries_; +}; + +// Repack kernels are additionally keyed by the activation (vec_dot_type) +// they were interleaved against, and by the interleave geometry -- carried +// alongside the ggml_type key as a small identifying suffix (e.g. "4x4", +// "8x8") so multiple repack variants can coexist for the same base type. +struct RepackKey { + ggml_type base_type; + ggml_type act_type; + int inter_size; + int nb_cols; + std::string label; // e.g. "q4_0_4x4_q8_0", used for display and as a stable map key + + bool operator<(const RepackKey &other) const { + return label < other.label; + } +}; + +template +class RepackRegistry { +public: + void register_reference(const RepackKey &key, std::string name, Fn fn) { + entries_[key.label].key = key; + entries_[key.label].reference = Impl{std::move(name), fn}; + } + void register_candidate(const RepackKey &key, std::string name, Fn fn) { + entries_[key.label].key = key; + entries_[key.label].candidates.push_back(Impl{std::move(name), fn}); + } + const Impl *reference(const std::string &label) const { + auto it = entries_.find(label); + if (it == entries_.end() || !it->second.reference.has_value()) { + return nullptr; + } + return &*it->second.reference; + } + const std::vector> &candidates(const std::string &label) const { + static const std::vector> empty; + auto it = entries_.find(label); + return it == entries_.end() ? empty : it->second.candidates; + } + std::vector keys() const { + std::vector out; + for (const auto &[label, entry] : entries_) { + if (entry.reference.has_value()) { + out.push_back(entry.key); + } + } + return out; + } + +private: + struct Entry { + RepackKey key{}; + std::optional> reference; + std::vector> candidates; + }; + std::map entries_; +}; + +struct KernelRegistries { + Registry quantize; + Registry dequantize; + Registry vec_dot; + RepackRegistry repack_quantize_mat; + RepackRegistry repack_gemv; + RepackRegistry repack_gemm; +}; diff --git a/apps/ggml/providers/README.md b/apps/ggml/providers/README.md new file mode 100644 index 000000000000..279d8bdbf775 --- /dev/null +++ b/apps/ggml/providers/README.md @@ -0,0 +1,60 @@ +# Adding a provider + +A "provider" is anything that registers one or more implementations into a +`KernelRegistries` (see `include/kernel_registry.h`). `ggml_provider.cpp` is the +provider shipped today; it is the *only* file that knows GGML's internal symbols +exist. A new provider -- e.g. your own from-scratch reference implementation, +starting with dequantize per the project's stated goal -- is just another +translation unit that does the same thing. + +## Steps + +1. Create `providers/_provider.h` declaring one function: + + ```cpp + void register__provider(KernelRegistries & registries); + ``` + +2. Create `providers/_provider.cpp` implementing it. For each + `(ggml_type, implementation)` pair you want benchmarked, call: + + ```cpp + registries.dequantize.register_candidate(GGML_TYPE_Q4_0, "my-dequant", my_dequantize_q4_0); + ``` + + matching the category's function-pointer typedef from `kernel_registry.h`: + + | category | typedef | signature | + | --------------------------------- | ----------------- | --------------------------------------------------------------------------------------------- | + | `quantize`, `repack_quantize_mat` | `quantize_fn_t` | `(const float* x, void* y, int64_t k)` | + | `dequantize` | `dequantize_fn_t` | `(const void* x, float* y, int64_t k)` | + | `vec_dot` | `vec_dot_fn_t` | `(int n, float* s, size_t bs, const void* vx, size_t bx, const void* vy, size_t by, int nrc)` | + | `repack_gemv`, `repack_gemm` | `gemx_fn_t` | `(int n, float* s, size_t bs, const void* vx, const void* vy, int nr, int nc)` | + + Use `register_candidate()` if GGML's existing reference should remain the + correctness ground truth for that type (the common case: you want to see + whether your implementation agrees with GGML and how fast it is). Use + `register_reference()` instead if your implementation should *become* the new + ground truth other candidates are compared against for that type -- the + harness doesn't care which provider a reference comes from. + +3. Add one line to `src/main.cpp`: + + ```cpp + register__provider(registries); + ``` + + next to the existing `register_ggml_provider(registries);` call. Nothing else + changes -- `bench_*.cpp`, the CLI, and the reporting code iterate whatever + ends up in the registries and don't know or care how many providers + contributed to them. + +## Repack keys + +`repack_quantize_mat`/`repack_gemv`/`repack_gemm` are keyed by `RepackKey` (base +type, activation type, interleave geometry, label string) rather than by +`ggml_type` alone, since several interleaved weight layouts can exist for the +same base type (e.g. `q4_0_4x4_q8_0` vs `q4_0_8x8_q8_0`). Reuse the `RepackKey` +values already registered by `ggml_provider.cpp` (see `k_repack_entries` in +`ggml_provider.cpp`) if you're providing an alternative gemv/gemm for an +existing layout; define your own `RepackKey` if you're introducing a new one. diff --git a/apps/ggml/providers/ggml_internal_abi.h b/apps/ggml/providers/ggml_internal_abi.h new file mode 100644 index 000000000000..a53ebf1bd905 --- /dev/null +++ b/apps/ggml/providers/ggml_internal_abi.h @@ -0,0 +1,295 @@ +#pragma once + +// PRIVATE, VERSION-PINNED ABI SURFACE -- READ BEFORE TOUCHING +// +// The functions declared below are internal implementation details of +// ggml-cpu (declared in the *uninstalled* headers src/ggml-cpu/quants.h and +// src/ggml-cpu/repack.h). GGML does not install those headers, does not +// document these symbols, does not version them, and offers no ABI +// stability guarantee for them whatsoever. +// +// They are reachable from an external application ONLY because ggml-cpu is +// built without -fvisibility=hidden: every plain, non-static C function ends +// up with default (exported) linker visibility by accident of the build +// configuration, not by design. This header is a hand-copied snapshot of +// the declarations in ggml (as of the commit this file was written against; +// see README.md) -- if a future GGML release renames, removes, or changes +// the signature of one of these functions, this header (and only this +// header + ggml_provider.cpp) will need updating. No other part of +// kernel-bench depends on GGML internals. +// +// Why we need this at all: GGML's public API (ggml_get_type_traits / +// ggml_get_type_traits_cpu, see include/ggml.h and include/ggml-cpu.h) +// exposes exactly one "reference" and one "dispatched" implementation per +// type for quantize/dequantize, which is enough for those two categories +// without touching anything private (see ggml_provider.cpp). It does NOT +// expose the always-available pure-C fallback for vec_dot or for the repack +// quantize_mat/gemv/gemm kernels -- the only way to reach those, and thus +// the only way to compare them against the (possibly arch-optimized) +// canonical symbol, is by declaring both names ourselves and letting the +// linker resolve them. +// +// IMPORTANT -- the `_generic` name does not always exist as its own link +// symbol: src/ggml-cpu/arch-fallback.h #defines it onto the canonical name +// -- as a textual macro substitution inside GGML's own .c/.cpp files -- for +// whichever functions the current architecture has no distinct optimized +// version of, and there is then only one function in the binary. (A weak +// C++ declaration doesn't paper over this: verified empirically that +// Darwin's ld64 hard-fails on an undefined `weak_import` symbol that has +// zero definitions anywhere in the link, and GNU ld's "resolve undefined +// weak to null" behavior isn't something to rely on portably either.) So +// this header mirrors arch-fallback.h's own collapsing, using the same +// preprocessor guards, restricted to the subset of functions declared +// below. When a `_generic` name collapses onto its canonical counterpart +// here exactly as it does inside GGML itself, ggml_provider.cpp's +// pointer-equality check (`generic_fn == canonical_fn`) naturally detects +// "single implementation, nothing to compare" for that kernel on this +// architecture. If GGML's own arch-fallback.h changes its collapsing list, +// this block needs updating to match; this is the one part of this file +// most likely to need attention when moving to a newer GGML. +#if defined(__aarch64__) || defined(__arm__) || defined(_M_ARM) || defined(_M_ARM64) +#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4 +#define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8 +#define ggml_gemv_iq4_nl_8x8_q8_0_generic ggml_gemv_iq4_nl_8x8_q8_0 +#define ggml_gemv_mxfp4_8x8_q8_0_generic ggml_gemv_mxfp4_8x8_q8_0 +#define ggml_gemv_q2_K_8x8_q8_K_generic ggml_gemv_q2_K_8x8_q8_K +#define ggml_gemm_iq4_nl_8x8_q8_0_generic ggml_gemm_iq4_nl_8x8_q8_0 +#define ggml_gemm_mxfp4_8x8_q8_0_generic ggml_gemm_mxfp4_8x8_q8_0 +#define ggml_gemm_q2_K_8x8_q8_K_generic ggml_gemm_q2_K_8x8_q8_K +#elif defined(__x86_64__) || defined(__i386__) || defined(_M_IX86) || defined(_M_X64) +#define ggml_vec_dot_nvfp4_q8_0_generic ggml_vec_dot_nvfp4_q8_0 +#define ggml_quantize_mat_q8_0_4x4_generic ggml_quantize_mat_q8_0_4x4 +#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4 +#define ggml_gemv_q4_0_4x4_q8_0_generic ggml_gemv_q4_0_4x4_q8_0 +#define ggml_gemv_q4_0_4x8_q8_0_generic ggml_gemv_q4_0_4x8_q8_0 +#define ggml_gemv_q4_K_8x4_q8_K_generic ggml_gemv_q4_K_8x4_q8_K +#define ggml_gemv_q5_K_8x4_q8_K_generic ggml_gemv_q5_K_8x4_q8_K +#define ggml_gemv_q5_K_8x8_q8_K_generic ggml_gemv_q5_K_8x8_q8_K +#define ggml_gemv_q6_K_8x4_q8_K_generic ggml_gemv_q6_K_8x4_q8_K +#define ggml_gemv_q6_K_8x8_q8_K_generic ggml_gemv_q6_K_8x8_q8_K +#define ggml_gemv_iq4_nl_4x4_q8_0_generic ggml_gemv_iq4_nl_4x4_q8_0 +#define ggml_gemv_mxfp4_4x4_q8_0_generic ggml_gemv_mxfp4_4x4_q8_0 +#define ggml_gemv_q8_0_4x4_q8_0_generic ggml_gemv_q8_0_4x4_q8_0 +#define ggml_gemv_q8_0_4x8_q8_0_generic ggml_gemv_q8_0_4x8_q8_0 +#define ggml_gemm_q4_0_4x4_q8_0_generic ggml_gemm_q4_0_4x4_q8_0 +#define ggml_gemm_q4_0_4x8_q8_0_generic ggml_gemm_q4_0_4x8_q8_0 +#define ggml_gemm_q4_K_8x4_q8_K_generic ggml_gemm_q4_K_8x4_q8_K +#define ggml_gemm_q5_K_8x4_q8_K_generic ggml_gemm_q5_K_8x4_q8_K +#define ggml_gemm_q5_K_8x8_q8_K_generic ggml_gemm_q5_K_8x8_q8_K +#define ggml_gemm_q6_K_8x4_q8_K_generic ggml_gemm_q6_K_8x4_q8_K +#define ggml_gemm_q6_K_8x8_q8_K_generic ggml_gemm_q6_K_8x8_q8_K +#define ggml_gemm_iq4_nl_4x4_q8_0_generic ggml_gemm_iq4_nl_4x4_q8_0 +#define ggml_gemm_mxfp4_4x4_q8_0_generic ggml_gemm_mxfp4_4x4_q8_0 +#define ggml_gemm_q8_0_4x4_q8_0_generic ggml_gemm_q8_0_4x4_q8_0 +#define ggml_gemm_q8_0_4x8_q8_0_generic ggml_gemm_q8_0_4x8_q8_0 +#elif defined(__POWERPC__) || defined(__powerpc__) || defined(__loongarch64) || defined(__riscv) || \ + defined(__s390x__) || defined(__wasm__) +// PowerPC/LoongArch/RISC-V/s390x/wasm each collapse a large, differently-shaped +// subset of quants.c/repack.cpp symbols (see arch-fallback.h) -- rather than +// transcribe five more per-arch lists by hand, collapse everything this +// header declares on these architectures. This is conservative in the safe +// direction: on an arch that actually kept a real optimized/generic split +// for some function, this makes that pairing look like "single +// implementation" instead of reporting it, but it will never misreport two +// genuinely different implementations as identical, and it will never +// produce a link error. +#define quantize_row_q8_K_generic quantize_row_q8_K +#define ggml_vec_dot_q1_0_q8_0_generic ggml_vec_dot_q1_0_q8_0 +#define ggml_vec_dot_mxfp4_q8_0_generic ggml_vec_dot_mxfp4_q8_0 +#define ggml_vec_dot_nvfp4_q8_0_generic ggml_vec_dot_nvfp4_q8_0 +#define ggml_vec_dot_tq1_0_q8_K_generic ggml_vec_dot_tq1_0_q8_K +#define ggml_vec_dot_tq2_0_q8_K_generic ggml_vec_dot_tq2_0_q8_K +#define ggml_vec_dot_q2_K_q8_K_generic ggml_vec_dot_q2_K_q8_K +#define ggml_vec_dot_iq2_xxs_q8_K_generic ggml_vec_dot_iq2_xxs_q8_K +#define ggml_vec_dot_iq2_xs_q8_K_generic ggml_vec_dot_iq2_xs_q8_K +#define ggml_vec_dot_iq2_s_q8_K_generic ggml_vec_dot_iq2_s_q8_K +#define ggml_vec_dot_iq3_xxs_q8_K_generic ggml_vec_dot_iq3_xxs_q8_K +#define ggml_vec_dot_iq3_s_q8_K_generic ggml_vec_dot_iq3_s_q8_K +#define ggml_vec_dot_iq1_s_q8_K_generic ggml_vec_dot_iq1_s_q8_K +#define ggml_vec_dot_iq1_m_q8_K_generic ggml_vec_dot_iq1_m_q8_K +#define ggml_vec_dot_iq4_nl_q8_0_generic ggml_vec_dot_iq4_nl_q8_0 +#define ggml_vec_dot_iq4_xs_q8_K_generic ggml_vec_dot_iq4_xs_q8_K +#define ggml_quantize_mat_q8_0_4x4_generic ggml_quantize_mat_q8_0_4x4 +#define ggml_quantize_mat_q8_0_4x8_generic ggml_quantize_mat_q8_0_4x8 +#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4 +#define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8 +#define ggml_gemv_q4_0_4x4_q8_0_generic ggml_gemv_q4_0_4x4_q8_0 +#define ggml_gemv_q4_0_4x8_q8_0_generic ggml_gemv_q4_0_4x8_q8_0 +#define ggml_gemv_q4_0_8x8_q8_0_generic ggml_gemv_q4_0_8x8_q8_0 +#define ggml_gemv_q2_K_8x8_q8_K_generic ggml_gemv_q2_K_8x8_q8_K +#define ggml_gemv_q4_K_8x4_q8_K_generic ggml_gemv_q4_K_8x4_q8_K +#define ggml_gemv_q4_K_8x8_q8_K_generic ggml_gemv_q4_K_8x8_q8_K +#define ggml_gemv_q5_K_8x4_q8_K_generic ggml_gemv_q5_K_8x4_q8_K +#define ggml_gemv_q5_K_8x8_q8_K_generic ggml_gemv_q5_K_8x8_q8_K +#define ggml_gemv_q6_K_8x4_q8_K_generic ggml_gemv_q6_K_8x4_q8_K +#define ggml_gemv_q6_K_8x8_q8_K_generic ggml_gemv_q6_K_8x8_q8_K +#define ggml_gemv_iq4_nl_4x4_q8_0_generic ggml_gemv_iq4_nl_4x4_q8_0 +#define ggml_gemv_iq4_nl_8x8_q8_0_generic ggml_gemv_iq4_nl_8x8_q8_0 +#define ggml_gemv_mxfp4_4x4_q8_0_generic ggml_gemv_mxfp4_4x4_q8_0 +#define ggml_gemv_mxfp4_8x8_q8_0_generic ggml_gemv_mxfp4_8x8_q8_0 +#define ggml_gemv_q8_0_4x4_q8_0_generic ggml_gemv_q8_0_4x4_q8_0 +#define ggml_gemv_q8_0_4x8_q8_0_generic ggml_gemv_q8_0_4x8_q8_0 +#define ggml_gemm_q4_0_4x4_q8_0_generic ggml_gemm_q4_0_4x4_q8_0 +#define ggml_gemm_q4_0_4x8_q8_0_generic ggml_gemm_q4_0_4x8_q8_0 +#define ggml_gemm_q4_0_8x8_q8_0_generic ggml_gemm_q4_0_8x8_q8_0 +#define ggml_gemm_q2_K_8x8_q8_K_generic ggml_gemm_q2_K_8x8_q8_K +#define ggml_gemm_q4_K_8x4_q8_K_generic ggml_gemm_q4_K_8x4_q8_K +#define ggml_gemm_q4_K_8x8_q8_K_generic ggml_gemm_q4_K_8x8_q8_K +#define ggml_gemm_q5_K_8x4_q8_K_generic ggml_gemm_q5_K_8x4_q8_K +#define ggml_gemm_q5_K_8x8_q8_K_generic ggml_gemm_q5_K_8x8_q8_K +#define ggml_gemm_q6_K_8x4_q8_K_generic ggml_gemm_q6_K_8x4_q8_K +#define ggml_gemm_q6_K_8x8_q8_K_generic ggml_gemm_q6_K_8x8_q8_K +#define ggml_gemm_iq4_nl_4x4_q8_0_generic ggml_gemm_iq4_nl_4x4_q8_0 +#define ggml_gemm_iq4_nl_8x8_q8_0_generic ggml_gemm_iq4_nl_8x8_q8_0 +#define ggml_gemm_mxfp4_4x4_q8_0_generic ggml_gemm_mxfp4_4x4_q8_0 +#define ggml_gemm_mxfp4_8x8_q8_0_generic ggml_gemm_mxfp4_8x8_q8_0 +#define ggml_gemm_q8_0_4x4_q8_0_generic ggml_gemm_q8_0_4x4_q8_0 +#define ggml_gemm_q8_0_4x8_q8_0_generic ggml_gemm_q8_0_4x8_q8_0 +#endif + +#include +#include + +#include // ggml_backend_buffer_type_t +#include // GGML_RESTRICT + +// NOTE: declared with ordinary C++ (mangled) linkage, matching the real +// src/ggml-cpu/repack.h -- this one declaration sits *before* that header's +// `extern "C" { ... }` block, unlike every quantize_row_*/vec_dot_*/gemv/gemm +// declaration below. +ggml_backend_buffer_type_t ggml_backend_cpu_repack_buffer_type(void); + +extern "C" { + +// -- src/ggml-cpu/quants.h: pure-C reference quantizer for Q8_K. Unlike +// every other quantized type, GGML_TYPE_Q8_K has no `from_float_ref` in the +// public ggml_get_type_traits() table (src/ggml.c) at all -- Q8_K is purely +// an internal activation format for K-quant vec_dot/gemv/gemm, never a +// row-conversion target -- so it needs this private symbol as its +// reference; see the special case in ggml_provider.cpp. +void quantize_row_q8_K_generic(const float *GGML_RESTRICT x, void *GGML_RESTRICT y, int64_t k); + +// -- src/ggml-quants.h: whole-matrix quantizers for the importance-matrix- +// only codebook types (IQ2_XXS, IQ2_XS, IQ1_S, IQ1_M). Unlike every other +// quantized type, these have no `from_float_ref` in the public +// ggml_get_type_traits() table at all -- GGML only exposes them through +// this differently-shaped `(src, dst, nrows, n_per_row, imatrix)` signature +// (nrows/n_per_row instead of a flat element count k, and an optional +// importance-matrix pointer), used only by the model-quantization tool. +// Called here with nrows=1, n_per_row=k, imatrix=nullptr to get a plain +// per-row reference, matching the shape every other type's from_float_ref +// already has; see the special case in ggml_provider.cpp. Declared to +// return size_t per GGML's real signature (the number of bytes written). +size_t quantize_iq2_xxs(const float *GGML_RESTRICT src, void *GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float *GGML_RESTRICT imatrix); +size_t quantize_iq2_xs(const float *GGML_RESTRICT src, void *GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float *GGML_RESTRICT imatrix); +size_t quantize_iq1_s(const float *GGML_RESTRICT src, void *GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float *GGML_RESTRICT imatrix); +size_t quantize_iq1_m(const float *GGML_RESTRICT src, void *GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float *GGML_RESTRICT imatrix); + +// -- src/ggml-cpu/quants.h: vec_dot, pure-C reference (collapsed onto the +// canonical, arch-dispatched symbol above on architectures with no distinct +// optimized implementation for that type) -- +void ggml_vec_dot_q1_0_q8_0_generic(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_q4_0_q8_0_generic(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_q4_1_q8_1_generic(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_q5_0_q8_0_generic(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_q5_1_q8_1_generic(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_q8_0_q8_0_generic(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_mxfp4_q8_0_generic(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_nvfp4_q8_0_generic(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_tq1_0_q8_K_generic(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_tq2_0_q8_K_generic(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_q2_K_q8_K_generic(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_q3_K_q8_K_generic(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_q4_K_q8_K_generic(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_q5_K_q8_K_generic(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_q6_K_q8_K_generic(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_iq2_xxs_q8_K_generic(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_iq2_xs_q8_K_generic(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_iq2_s_q8_K_generic(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_iq3_xxs_q8_K_generic(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_iq3_s_q8_K_generic(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_iq1_s_q8_K_generic(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_iq1_m_q8_K_generic(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_iq4_nl_q8_0_generic(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_iq4_xs_q8_K_generic(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); + +// -- src/ggml-cpu/repack.h: activation packing (float -> interleaved q8 blocks), canonical + reference -- +void ggml_quantize_mat_q8_0_4x4(const float *GGML_RESTRICT x, void *GGML_RESTRICT vy, int64_t k); +void ggml_quantize_mat_q8_0_4x8(const float *GGML_RESTRICT x, void *GGML_RESTRICT vy, int64_t k); +void ggml_quantize_mat_q8_K_4x4(const float *GGML_RESTRICT x, void *GGML_RESTRICT vy, int64_t k); +void ggml_quantize_mat_q8_K_4x8(const float *GGML_RESTRICT x, void *GGML_RESTRICT vy, int64_t k); +void ggml_quantize_mat_q8_0_4x4_generic(const float *GGML_RESTRICT x, void *GGML_RESTRICT vy, int64_t k); +void ggml_quantize_mat_q8_0_4x8_generic(const float *GGML_RESTRICT x, void *GGML_RESTRICT vy, int64_t k); +void ggml_quantize_mat_q8_K_4x4_generic(const float *GGML_RESTRICT x, void *GGML_RESTRICT vy, int64_t k); +void ggml_quantize_mat_q8_K_4x8_generic(const float *GGML_RESTRICT x, void *GGML_RESTRICT vy, int64_t k); + +// -- src/ggml-cpu/repack.h: gemv/gemm over packed weight blocks, canonical + reference -- +void ggml_gemv_q4_0_4x4_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q4_0_4x8_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q4_0_8x8_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q2_K_8x8_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q4_K_8x4_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q4_K_8x8_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q5_K_8x4_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q5_K_8x8_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q6_K_8x4_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q6_K_8x8_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_iq4_nl_4x4_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_iq4_nl_8x8_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_mxfp4_4x4_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_mxfp4_8x8_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q8_0_4x4_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q8_0_4x8_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); + +void ggml_gemm_q4_0_4x4_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q4_0_4x8_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q4_0_8x8_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q2_K_8x8_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q4_K_8x4_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q4_K_8x8_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q5_K_8x4_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q5_K_8x8_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q6_K_8x4_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q6_K_8x8_q8_K(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_iq4_nl_4x4_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_iq4_nl_8x8_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_mxfp4_4x4_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_mxfp4_8x8_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q8_0_4x4_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q8_0_4x8_q8_0(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); + +void ggml_gemv_q4_0_4x4_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q4_0_4x8_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q4_0_8x8_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q2_K_8x8_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q4_K_8x4_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q4_K_8x8_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q5_K_8x4_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q5_K_8x8_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q6_K_8x4_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q6_K_8x8_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_iq4_nl_4x4_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_iq4_nl_8x8_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_mxfp4_4x4_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_mxfp4_8x8_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q8_0_4x4_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemv_q8_0_4x8_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); + +void ggml_gemm_q4_0_4x4_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q4_0_4x8_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q4_0_8x8_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q2_K_8x8_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q4_K_8x4_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q4_K_8x8_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q5_K_8x4_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q5_K_8x8_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q6_K_8x4_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q6_K_8x8_q8_K_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_iq4_nl_4x4_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_iq4_nl_8x8_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_mxfp4_4x4_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_mxfp4_8x8_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q8_0_4x4_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); +void ggml_gemm_q8_0_4x8_q8_0_generic(int n, float *GGML_RESTRICT s, size_t bs, const void *GGML_RESTRICT vx, const void *GGML_RESTRICT vy, int nr, int nc); + +} // extern "C" diff --git a/apps/ggml/providers/ggml_provider.cpp b/apps/ggml/providers/ggml_provider.cpp new file mode 100644 index 000000000000..928233cdf13f --- /dev/null +++ b/apps/ggml/providers/ggml_provider.cpp @@ -0,0 +1,293 @@ +#include "ggml_provider.h" +#include "ggml_internal_abi.h" + +#include +#include + +#include + +namespace { + +// type -> pure-C reference vec_dot (src/ggml-cpu/quants.h `_generic` symbols). +// The canonical (possibly arch-optimized) candidate is obtained separately, +// through the PUBLIC ggml_get_type_traits_cpu(type)->vec_dot. +struct VecDotRef { + ggml_type type; + vec_dot_fn_t fn; +}; + +const VecDotRef k_vec_dot_refs[] = { + {GGML_TYPE_Q1_0, ggml_vec_dot_q1_0_q8_0_generic}, + {GGML_TYPE_Q4_0, ggml_vec_dot_q4_0_q8_0_generic}, + {GGML_TYPE_Q4_1, ggml_vec_dot_q4_1_q8_1_generic}, + {GGML_TYPE_Q5_0, ggml_vec_dot_q5_0_q8_0_generic}, + {GGML_TYPE_Q5_1, ggml_vec_dot_q5_1_q8_1_generic}, + {GGML_TYPE_Q8_0, ggml_vec_dot_q8_0_q8_0_generic}, + {GGML_TYPE_MXFP4, ggml_vec_dot_mxfp4_q8_0_generic}, + {GGML_TYPE_NVFP4, ggml_vec_dot_nvfp4_q8_0_generic}, + {GGML_TYPE_Q2_K, ggml_vec_dot_q2_K_q8_K_generic}, + {GGML_TYPE_Q3_K, ggml_vec_dot_q3_K_q8_K_generic}, + {GGML_TYPE_Q4_K, ggml_vec_dot_q4_K_q8_K_generic}, + {GGML_TYPE_Q5_K, ggml_vec_dot_q5_K_q8_K_generic}, + {GGML_TYPE_Q6_K, ggml_vec_dot_q6_K_q8_K_generic}, + {GGML_TYPE_TQ1_0, ggml_vec_dot_tq1_0_q8_K_generic}, + {GGML_TYPE_TQ2_0, ggml_vec_dot_tq2_0_q8_K_generic}, + {GGML_TYPE_IQ2_XXS, ggml_vec_dot_iq2_xxs_q8_K_generic}, + {GGML_TYPE_IQ2_XS, ggml_vec_dot_iq2_xs_q8_K_generic}, + {GGML_TYPE_IQ2_S, ggml_vec_dot_iq2_s_q8_K_generic}, + {GGML_TYPE_IQ3_XXS, ggml_vec_dot_iq3_xxs_q8_K_generic}, + {GGML_TYPE_IQ3_S, ggml_vec_dot_iq3_s_q8_K_generic}, + {GGML_TYPE_IQ1_S, ggml_vec_dot_iq1_s_q8_K_generic}, + {GGML_TYPE_IQ1_M, ggml_vec_dot_iq1_m_q8_K_generic}, + {GGML_TYPE_IQ4_NL, ggml_vec_dot_iq4_nl_q8_0_generic}, + {GGML_TYPE_IQ4_XS, ggml_vec_dot_iq4_xs_q8_K_generic}, +}; + +// The 9 repack combinations enumerated in +// ggml_repack_get_optimal_repack_type() (src/ggml-cpu/repack.cpp:4528-4560). +struct RepackEntry { + RepackKey key; + quantize_fn_t quantize_mat; + quantize_fn_t quantize_mat_generic; + gemx_fn_t gemv; + gemx_fn_t gemv_generic; + gemx_fn_t gemm; + gemx_fn_t gemm_generic; +}; + +const RepackEntry k_repack_entries[] = { + {{GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 4, 4, "q4_0_4x4_q8_0"}, + ggml_quantize_mat_q8_0_4x4, + ggml_quantize_mat_q8_0_4x4_generic, + ggml_gemv_q4_0_4x4_q8_0, + ggml_gemv_q4_0_4x4_q8_0_generic, + ggml_gemm_q4_0_4x4_q8_0, + ggml_gemm_q4_0_4x4_q8_0_generic}, + {{GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 8, 4, "q4_0_4x8_q8_0"}, + ggml_quantize_mat_q8_0_4x8, + ggml_quantize_mat_q8_0_4x8_generic, + ggml_gemv_q4_0_4x8_q8_0, + ggml_gemv_q4_0_4x8_q8_0_generic, + ggml_gemm_q4_0_4x8_q8_0, + ggml_gemm_q4_0_4x8_q8_0_generic}, + {{GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 8, 8, "q4_0_8x8_q8_0"}, + ggml_quantize_mat_q8_0_4x8, + ggml_quantize_mat_q8_0_4x8_generic, + ggml_gemv_q4_0_8x8_q8_0, + ggml_gemv_q4_0_8x8_q8_0_generic, + ggml_gemm_q4_0_8x8_q8_0, + ggml_gemm_q4_0_8x8_q8_0_generic}, + {{GGML_TYPE_Q4_K, GGML_TYPE_Q8_K, 8, 4, "q4_K_8x4_q8_K"}, + ggml_quantize_mat_q8_K_4x4, + ggml_quantize_mat_q8_K_4x4_generic, + ggml_gemv_q4_K_8x4_q8_K, + ggml_gemv_q4_K_8x4_q8_K_generic, + ggml_gemm_q4_K_8x4_q8_K, + ggml_gemm_q4_K_8x4_q8_K_generic}, + {{GGML_TYPE_Q4_K, GGML_TYPE_Q8_K, 8, 8, "q4_K_8x8_q8_K"}, + ggml_quantize_mat_q8_K_4x8, + ggml_quantize_mat_q8_K_4x8_generic, + ggml_gemv_q4_K_8x8_q8_K, + ggml_gemv_q4_K_8x8_q8_K_generic, + ggml_gemm_q4_K_8x8_q8_K, + ggml_gemm_q4_K_8x8_q8_K_generic}, + {{GGML_TYPE_Q5_K, GGML_TYPE_Q8_K, 8, 4, "q5_K_8x4_q8_K"}, + ggml_quantize_mat_q8_K_4x4, + ggml_quantize_mat_q8_K_4x4_generic, + ggml_gemv_q5_K_8x4_q8_K, + ggml_gemv_q5_K_8x4_q8_K_generic, + ggml_gemm_q5_K_8x4_q8_K, + ggml_gemm_q5_K_8x4_q8_K_generic}, + {{GGML_TYPE_Q5_K, GGML_TYPE_Q8_K, 8, 8, "q5_K_8x8_q8_K"}, + ggml_quantize_mat_q8_K_4x8, + ggml_quantize_mat_q8_K_4x8_generic, + ggml_gemv_q5_K_8x8_q8_K, + ggml_gemv_q5_K_8x8_q8_K_generic, + ggml_gemm_q5_K_8x8_q8_K, + ggml_gemm_q5_K_8x8_q8_K_generic}, + {{GGML_TYPE_Q6_K, GGML_TYPE_Q8_K, 8, 4, "q6_K_8x4_q8_K"}, + ggml_quantize_mat_q8_K_4x4, + ggml_quantize_mat_q8_K_4x4_generic, + ggml_gemv_q6_K_8x4_q8_K, + ggml_gemv_q6_K_8x4_q8_K_generic, + ggml_gemm_q6_K_8x4_q8_K, + ggml_gemm_q6_K_8x4_q8_K_generic}, + {{GGML_TYPE_Q6_K, GGML_TYPE_Q8_K, 8, 8, "q6_K_8x8_q8_K"}, + ggml_quantize_mat_q8_K_4x8, + ggml_quantize_mat_q8_K_4x8_generic, + ggml_gemv_q6_K_8x8_q8_K, + ggml_gemv_q6_K_8x8_q8_K_generic, + ggml_gemm_q6_K_8x8_q8_K, + ggml_gemm_q6_K_8x8_q8_K_generic}, + {{GGML_TYPE_Q2_K, GGML_TYPE_Q8_K, 8, 8, "q2_K_8x8_q8_K"}, + ggml_quantize_mat_q8_K_4x8, + ggml_quantize_mat_q8_K_4x8_generic, + ggml_gemv_q2_K_8x8_q8_K, + ggml_gemv_q2_K_8x8_q8_K_generic, + ggml_gemm_q2_K_8x8_q8_K, + ggml_gemm_q2_K_8x8_q8_K_generic}, + {{GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_0, 4, 4, "iq4_nl_4x4_q8_0"}, + ggml_quantize_mat_q8_0_4x4, + ggml_quantize_mat_q8_0_4x4_generic, + ggml_gemv_iq4_nl_4x4_q8_0, + ggml_gemv_iq4_nl_4x4_q8_0_generic, + ggml_gemm_iq4_nl_4x4_q8_0, + ggml_gemm_iq4_nl_4x4_q8_0_generic}, + {{GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_0, 8, 8, "iq4_nl_8x8_q8_0"}, + ggml_quantize_mat_q8_0_4x8, + ggml_quantize_mat_q8_0_4x8_generic, + ggml_gemv_iq4_nl_8x8_q8_0, + ggml_gemv_iq4_nl_8x8_q8_0_generic, + ggml_gemm_iq4_nl_8x8_q8_0, + ggml_gemm_iq4_nl_8x8_q8_0_generic}, + {{GGML_TYPE_MXFP4, GGML_TYPE_Q8_0, 4, 4, "mxfp4_4x4_q8_0"}, + ggml_quantize_mat_q8_0_4x4, + ggml_quantize_mat_q8_0_4x4_generic, + ggml_gemv_mxfp4_4x4_q8_0, + ggml_gemv_mxfp4_4x4_q8_0_generic, + ggml_gemm_mxfp4_4x4_q8_0, + ggml_gemm_mxfp4_4x4_q8_0_generic}, + {{GGML_TYPE_MXFP4, GGML_TYPE_Q8_0, 8, 8, "mxfp4_8x8_q8_0"}, + ggml_quantize_mat_q8_0_4x8, + ggml_quantize_mat_q8_0_4x8_generic, + ggml_gemv_mxfp4_8x8_q8_0, + ggml_gemv_mxfp4_8x8_q8_0_generic, + ggml_gemm_mxfp4_8x8_q8_0, + ggml_gemm_mxfp4_8x8_q8_0_generic}, + {{GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, 4, 4, "q8_0_4x4_q8_0"}, + ggml_quantize_mat_q8_0_4x4, + ggml_quantize_mat_q8_0_4x4_generic, + ggml_gemv_q8_0_4x4_q8_0, + ggml_gemv_q8_0_4x4_q8_0_generic, + ggml_gemm_q8_0_4x4_q8_0, + ggml_gemm_q8_0_4x4_q8_0_generic}, + {{GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, 8, 4, "q8_0_4x8_q8_0"}, + ggml_quantize_mat_q8_0_4x8, + ggml_quantize_mat_q8_0_4x8_generic, + ggml_gemv_q8_0_4x8_q8_0, + ggml_gemv_q8_0_4x8_q8_0_generic, + ggml_gemm_q8_0_4x8_q8_0, + ggml_gemm_q8_0_4x8_q8_0_generic}, +}; + +// Thin adapters from GGML's whole-matrix `(src, dst, nrows, n_per_row, +// imatrix)` quantizer signature (see ggml_internal_abi.h) down to +// quantize_fn_t's flat `(x, y, k)` shape, matching what every other type's +// from_float_ref already looks like: one row, no importance weighting. All +// four of these implementations actually require a non-null quant_weights +// pointer (a real GGML_ASSERT for IQ2_XXS/IQ2_XS/IQ1_S; commented out but +// still exercised for IQ1_M) -- a uniform (all-1.0) weighting is passed so +// every element is treated as equally important, the closest equivalent to +// "no importance weighting" these quantizers support. +void quantize_iq2_xxs_row(const float *x, void *y, int64_t k) { + std::vector w(k, 1.0f); + quantize_iq2_xxs(x, y, 1, k, w.data()); +} +void quantize_iq2_xs_row(const float *x, void *y, int64_t k) { + std::vector w(k, 1.0f); + quantize_iq2_xs(x, y, 1, k, w.data()); +} +void quantize_iq1_s_row(const float *x, void *y, int64_t k) { + std::vector w(k, 1.0f); + quantize_iq1_s(x, y, 1, k, w.data()); +} +void quantize_iq1_m_row(const float *x, void *y, int64_t k) { + std::vector w(k, 1.0f); + quantize_iq1_m(x, y, 1, k, w.data()); +} + +} // namespace + +void register_ggml_provider(KernelRegistries ®istries) { + ggml_cpu_init(); + + // -- quantize / dequantize: fully public API -- + for (int t = 0; t < GGML_TYPE_COUNT; ++t) { + const ggml_type type = static_cast(t); + const ggml_type_traits *tt = ggml_get_type_traits(type); + const ggml_type_traits_cpu *tc = ggml_get_type_traits_cpu(type); + if (!tt || !tc) { + continue; + } + // Deliberately NOT calling ggml_quantize_init(type) here: for + // IQ2_XXS/IQ2_XS/IQ2_S/IQ1_S/IQ1_M/IQ3_XXS/IQ3_S it builds a nearest- + // neighbor lookup table via an O(43692 * grid_size) search with a + // qsort per row (src/ggml-quants.c, iq2xs_init_impl/iq3xs_init_impl) + // -- genuinely slow (hundreds of ms), and doing it here for all 42 + // types unconditionally at registration time means paying for it + // before any benchmark has printed a single row, even for runs + // (e.g. --repack) that never touch these types at all. Each bench_*.cpp + // calls it lazily, once, right before it first actually invokes one + // of these types' quantize/dequantize/vec_dot functions instead. + + if (tt->from_float_ref) { + registries.quantize.register_reference(type, "ggml-ref", tt->from_float_ref); + if (tc->from_float) { + registries.quantize.register_candidate(type, "ggml-cpu", tc->from_float); + } + } + if (tt->to_float) { + registries.dequantize.register_reference(type, "ggml-ref", tt->to_float); + // No candidate registered yet: GGML has exactly one dequantize + // implementation per type (src/ggml-quants.c, arch-independent). + // This is where a from-scratch dequantizer plugs in later. + } + } + + // GGML_TYPE_Q8_K has no public from_float_ref (see quantize_row_q8_K_generic's + // doc comment in ggml_internal_abi.h) -- it's the activation format for every + // K-quant vec_dot/gemv/gemm, so without this, all of those silently have no + // valid input to quantize into and get skipped by the benchmarks. + { + const ggml_type_traits_cpu *tc = ggml_get_type_traits_cpu(GGML_TYPE_Q8_K); + registries.quantize.register_reference(GGML_TYPE_Q8_K, "ggml-generic", quantize_row_q8_K_generic); + if (tc && tc->from_float) { + registries.quantize.register_candidate(GGML_TYPE_Q8_K, "ggml-cpu", tc->from_float); + } + } + + // GGML_TYPE_IQ2_XXS/IQ2_XS/IQ1_S/IQ1_M have no public from_float_ref + // either (see quantize_iq2_xxs's doc comment in ggml_internal_abi.h) -- + // they're importance-matrix-only codebook types whose only public + // quantizer takes a different, whole-matrix signature. Without this, + // these 4 types would never appear in the quantize/dequantize + // benchmarks at all (bench_dequantize.cpp requires both a quantize and + // a dequantize reference to exist before it will test a type). + registries.quantize.register_reference(GGML_TYPE_IQ2_XXS, "ggml-ref", quantize_iq2_xxs_row); + registries.quantize.register_reference(GGML_TYPE_IQ2_XS, "ggml-ref", quantize_iq2_xs_row); + registries.quantize.register_reference(GGML_TYPE_IQ1_S, "ggml-ref", quantize_iq1_s_row); + registries.quantize.register_reference(GGML_TYPE_IQ1_M, "ggml-ref", quantize_iq1_m_row); + + // arch-fallback.h #defines a `_generic` symbol onto its canonical + // counterpart -- inside GGML's own source files -- for whichever + // functions the current architecture has no distinct optimized version + // of, and which functions that applies to varies by architecture. The + // `_generic` declarations in ggml_internal_abi.h are marked + // GGML_BENCH_WEAK precisely so that case resolves to a null function + // pointer here instead of a link error: when null, there is only one + // real implementation, so it becomes the reference with no candidate, + // rather than fabricating a comparison against nothing. + auto register_pair = [](auto ®istry, const auto &key, auto generic_fn, auto canonical_fn) { + if (generic_fn) { + registry.register_reference(key, "ggml-generic", generic_fn); + if (canonical_fn) { + registry.register_candidate(key, "ggml-cpu", canonical_fn); + } + } else if (canonical_fn) { + registry.register_reference(key, "ggml-cpu", canonical_fn); + } + }; + + // -- vec_dot: public candidate, private (possibly weak-null) reference -- + for (const auto &ref : k_vec_dot_refs) { + const ggml_type_traits_cpu *tc = ggml_get_type_traits_cpu(ref.type); + register_pair(registries.vec_dot, ref.type, ref.fn, tc ? tc->vec_dot : nullptr); + } + + // -- repack: private reference and candidate (no public accessor exists) -- + for (const auto &e : k_repack_entries) { + register_pair(registries.repack_quantize_mat, e.key, e.quantize_mat_generic, e.quantize_mat); + register_pair(registries.repack_gemv, e.key, e.gemv_generic, e.gemv); + register_pair(registries.repack_gemm, e.key, e.gemm_generic, e.gemm); + } +} diff --git a/apps/ggml/providers/ggml_provider.h b/apps/ggml/providers/ggml_provider.h new file mode 100644 index 000000000000..bca2c7ec42fb --- /dev/null +++ b/apps/ggml/providers/ggml_provider.h @@ -0,0 +1,21 @@ +#pragma once + +#include "kernel_registry.h" + +// Registers GGML's own implementations into `registries`: +// +// quantize / dequantize -- entirely via GGML's PUBLIC API +// (ggml_get_type_traits / ggml_get_type_traits_cpu, see include/ggml.h and +// include/ggml-cpu.h): `from_float_ref`/`to_float` (GGML's own documented +// "reference" routines) become the Registry reference, and the +// CPU-dispatched `from_float` becomes a candidate. No private header used. +// +// vec_dot / repack (quantize_mat, gemv, gemm) -- these categories have no +// public way to reach the always-available pure-C fallback, so the +// `_generic`-suffixed symbol (declared in ggml_internal_abi.h) is +// registered as the reference and the canonical symbol (reached publicly +// for vec_dot via ggml_get_type_traits_cpu, and privately for repack, +// which has no public accessor at all) is registered as a candidate. +// +// This is the only file in kernel-bench that knows GGML exists. +void register_ggml_provider(KernelRegistries ®istries); diff --git a/apps/ggml/providers/halide_provider.cpp b/apps/ggml/providers/halide_provider.cpp new file mode 100644 index 000000000000..09df8bcc2333 --- /dev/null +++ b/apps/ggml/providers/halide_provider.cpp @@ -0,0 +1,236 @@ +#include "halide_provider.h" + +#include + +void register_halide_provider(KernelRegistries ®istries) { + registries.quantize.register_candidate(GGML_TYPE_Q4_0, "halide", ggml_quants_halide_quantize_q4_0); + registries.dequantize.register_candidate(GGML_TYPE_Q4_0, "halide", ggml_quants_halide_dequantize_q4_0); + registries.vec_dot.register_candidate(GGML_TYPE_Q4_0, "halide", ggml_quants_halide_vec_dot_q4_0_q8_0); + + registries.quantize.register_candidate(GGML_TYPE_Q4_1, "halide", ggml_quants_halide_quantize_q4_1); + registries.dequantize.register_candidate(GGML_TYPE_Q4_1, "halide", ggml_quants_halide_dequantize_q4_1); + registries.vec_dot.register_candidate(GGML_TYPE_Q4_1, "halide", ggml_quants_halide_vec_dot_q4_1_q8_1); + + registries.quantize.register_candidate(GGML_TYPE_Q5_0, "halide", ggml_quants_halide_quantize_q5_0); + registries.dequantize.register_candidate(GGML_TYPE_Q5_0, "halide", ggml_quants_halide_dequantize_q5_0); + registries.vec_dot.register_candidate(GGML_TYPE_Q5_0, "halide", ggml_quants_halide_vec_dot_q5_0_q8_0); + + registries.quantize.register_candidate(GGML_TYPE_Q5_1, "halide", ggml_quants_halide_quantize_q5_1); + registries.dequantize.register_candidate(GGML_TYPE_Q5_1, "halide", ggml_quants_halide_dequantize_q5_1); + registries.vec_dot.register_candidate(GGML_TYPE_Q5_1, "halide", ggml_quants_halide_vec_dot_q5_1_q8_1); + + registries.quantize.register_candidate(GGML_TYPE_Q8_0, "halide", ggml_quants_halide_quantize_q8_0); + registries.dequantize.register_candidate(GGML_TYPE_Q8_0, "halide", ggml_quants_halide_dequantize_q8_0); + registries.vec_dot.register_candidate(GGML_TYPE_Q8_0, "halide", ggml_quants_halide_vec_dot_q8_0_q8_0); + + // Q8_1 is activation-only (GGML has no public to_float for it, so + // ggml_provider.cpp registers no dequantize reference either -- the + // harness's bench_dequantize.cpp already skips types with no reference). + registries.quantize.register_candidate(GGML_TYPE_Q8_1, "halide", ggml_quants_halide_quantize_q8_1); + + // Q8_K is also activation-only, but unlike Q8_1, GGML doesn't even + // register a public from_float_ref for it -- ggml_provider.cpp's + // special case (using the private quantize_row_q8_K_generic) is what + // supplies the reference this candidate is compared against. + registries.quantize.register_candidate(GGML_TYPE_Q8_K, "halide", ggml_quants_halide_quantize_q8_k); + + // Q2_K, Q6_K: dequantize is a genuine from-scratch Halide candidate. + // Quantize is scaffolding that itself calls out to GGML's own + // reference (see halide/ggml_extern_quantize.cpp) pending a + // from-scratch port of GGML's iterative scale search -- it's still + // registered as a candidate so the harness's plumbing is exercised + // end-to-end, but it will trivially match (same underlying code path). + // vec_dot is a genuine from-scratch candidate (against Q8_K). + registries.quantize.register_candidate(GGML_TYPE_Q2_K, "halide", ggml_quants_halide_quantize_q2_k); + registries.dequantize.register_candidate(GGML_TYPE_Q2_K, "halide", ggml_quants_halide_dequantize_q2_k); + registries.vec_dot.register_candidate(GGML_TYPE_Q2_K, "halide", ggml_quants_halide_vec_dot_q2_k_q8_k); + + registries.quantize.register_candidate(GGML_TYPE_Q6_K, "halide", ggml_quants_halide_quantize_q6_k); + registries.dequantize.register_candidate(GGML_TYPE_Q6_K, "halide", ggml_quants_halide_dequantize_q6_k); + registries.vec_dot.register_candidate(GGML_TYPE_Q6_K, "halide", ggml_quants_halide_vec_dot_q6_k_q8_k); + + registries.quantize.register_candidate(GGML_TYPE_Q4_K, "halide", ggml_quants_halide_quantize_q4_k); + registries.dequantize.register_candidate(GGML_TYPE_Q4_K, "halide", ggml_quants_halide_dequantize_q4_k); + registries.vec_dot.register_candidate(GGML_TYPE_Q4_K, "halide", ggml_quants_halide_vec_dot_q4_k_q8_k); + + registries.quantize.register_candidate(GGML_TYPE_Q5_K, "halide", ggml_quants_halide_quantize_q5_k); + registries.dequantize.register_candidate(GGML_TYPE_Q5_K, "halide", ggml_quants_halide_dequantize_q5_k); + registries.vec_dot.register_candidate(GGML_TYPE_Q5_K, "halide", ggml_quants_halide_vec_dot_q5_k_q8_k); + + registries.quantize.register_candidate(GGML_TYPE_Q3_K, "halide", ggml_quants_halide_quantize_q3_k); + registries.dequantize.register_candidate(GGML_TYPE_Q3_K, "halide", ggml_quants_halide_dequantize_q3_k); + registries.vec_dot.register_candidate(GGML_TYPE_Q3_K, "halide", ggml_quants_halide_vec_dot_q3_k_q8_k); + + registries.quantize.register_candidate(GGML_TYPE_Q1_0, "halide", ggml_quants_halide_quantize_q1_0); + registries.dequantize.register_candidate(GGML_TYPE_Q1_0, "halide", ggml_quants_halide_dequantize_q1_0); + registries.vec_dot.register_candidate(GGML_TYPE_Q1_0, "halide", ggml_quants_halide_vec_dot_q1_0_q8_0); + + registries.quantize.register_candidate(GGML_TYPE_MXFP4, "halide", ggml_quants_halide_quantize_mxfp4); + registries.dequantize.register_candidate(GGML_TYPE_MXFP4, "halide", ggml_quants_halide_dequantize_mxfp4); + registries.vec_dot.register_candidate(GGML_TYPE_MXFP4, "halide", ggml_quants_halide_vec_dot_mxfp4_q8_0); + + registries.quantize.register_candidate(GGML_TYPE_NVFP4, "halide", ggml_quants_halide_quantize_nvfp4); + registries.dequantize.register_candidate(GGML_TYPE_NVFP4, "halide", ggml_quants_halide_dequantize_nvfp4); + registries.vec_dot.register_candidate(GGML_TYPE_NVFP4, "halide", ggml_quants_halide_vec_dot_nvfp4_q8_0); + + registries.quantize.register_candidate(GGML_TYPE_IQ4_NL, "halide", ggml_quants_halide_quantize_iq4_nl); + registries.dequantize.register_candidate(GGML_TYPE_IQ4_NL, "halide", ggml_quants_halide_dequantize_iq4_nl); + registries.vec_dot.register_candidate(GGML_TYPE_IQ4_NL, "halide", ggml_quants_halide_vec_dot_iq4_nl_q8_0); + + registries.quantize.register_candidate(GGML_TYPE_IQ4_XS, "halide", ggml_quants_halide_quantize_iq4_xs); + registries.dequantize.register_candidate(GGML_TYPE_IQ4_XS, "halide", ggml_quants_halide_dequantize_iq4_xs); + registries.vec_dot.register_candidate(GGML_TYPE_IQ4_XS, "halide", ggml_quants_halide_vec_dot_iq4_xs_q8_k); + + registries.quantize.register_candidate(GGML_TYPE_TQ1_0, "halide", ggml_quants_halide_quantize_tq1_0); + registries.dequantize.register_candidate(GGML_TYPE_TQ1_0, "halide", ggml_quants_halide_dequantize_tq1_0); + registries.vec_dot.register_candidate(GGML_TYPE_TQ1_0, "halide", ggml_quants_halide_vec_dot_tq1_0_q8_k); + + registries.quantize.register_candidate(GGML_TYPE_TQ2_0, "halide", ggml_quants_halide_quantize_tq2_0); + registries.dequantize.register_candidate(GGML_TYPE_TQ2_0, "halide", ggml_quants_halide_dequantize_tq2_0); + registries.vec_dot.register_candidate(GGML_TYPE_TQ2_0, "halide", ggml_quants_halide_vec_dot_tq2_0_q8_k); + + // IQ2_XXS: dequantize only (see ggml_quants.h for why); vec_dot is still + // a genuine from-scratch candidate (GGML's own reference quantizer is + // used to produce the test/benchmark input, via ggml_provider.cpp). + registries.dequantize.register_candidate(GGML_TYPE_IQ2_XXS, "halide", ggml_quants_halide_dequantize_iq2_xxs); + registries.vec_dot.register_candidate(GGML_TYPE_IQ2_XXS, "halide", ggml_quants_halide_vec_dot_iq2_xxs_q8_k); + + registries.dequantize.register_candidate(GGML_TYPE_IQ2_XS, "halide", ggml_quants_halide_dequantize_iq2_xs); + registries.vec_dot.register_candidate(GGML_TYPE_IQ2_XS, "halide", ggml_quants_halide_vec_dot_iq2_xs_q8_k); + + registries.quantize.register_candidate(GGML_TYPE_IQ2_S, "halide", ggml_quants_halide_quantize_iq2_s); + registries.dequantize.register_candidate(GGML_TYPE_IQ2_S, "halide", ggml_quants_halide_dequantize_iq2_s); + registries.vec_dot.register_candidate(GGML_TYPE_IQ2_S, "halide", ggml_quants_halide_vec_dot_iq2_s_q8_k); + + registries.quantize.register_candidate(GGML_TYPE_IQ3_XXS, "halide", ggml_quants_halide_quantize_iq3_xxs); + registries.dequantize.register_candidate(GGML_TYPE_IQ3_XXS, "halide", ggml_quants_halide_dequantize_iq3_xxs); + registries.vec_dot.register_candidate(GGML_TYPE_IQ3_XXS, "halide", ggml_quants_halide_vec_dot_iq3_xxs_q8_k); + + registries.quantize.register_candidate(GGML_TYPE_IQ3_S, "halide", ggml_quants_halide_quantize_iq3_s); + registries.dequantize.register_candidate(GGML_TYPE_IQ3_S, "halide", ggml_quants_halide_dequantize_iq3_s); + registries.vec_dot.register_candidate(GGML_TYPE_IQ3_S, "halide", ggml_quants_halide_vec_dot_iq3_s_q8_k); + + registries.dequantize.register_candidate(GGML_TYPE_IQ1_S, "halide", ggml_quants_halide_dequantize_iq1_s); + registries.vec_dot.register_candidate(GGML_TYPE_IQ1_S, "halide", ggml_quants_halide_vec_dot_iq1_s_q8_k); + + registries.dequantize.register_candidate(GGML_TYPE_IQ1_M, "halide", ggml_quants_halide_dequantize_iq1_m); + registries.vec_dot.register_candidate(GGML_TYPE_IQ1_M, "halide", ggml_quants_halide_vec_dot_iq1_m_q8_k); + + // F16, BF16: plain float casts, not "quantized" types, but both + // directions are fully native Halide. No vec_dot: not part of the + // quantized-format vec_dot sweep this directory otherwise covers. + registries.quantize.register_candidate(GGML_TYPE_F16, "halide", ggml_quants_halide_quantize_f16); + registries.dequantize.register_candidate(GGML_TYPE_F16, "halide", ggml_quants_halide_dequantize_f16); + + registries.quantize.register_candidate(GGML_TYPE_BF16, "halide", ggml_quants_halide_quantize_bf16); + registries.dequantize.register_candidate(GGML_TYPE_BF16, "halide", ggml_quants_halide_dequantize_bf16); + + // Repack quantize_mat: GGML itself only has 4 distinct implementations + // (2 activation formats x 2 interleave widths), reused across every + // repack weight type that shares one -- see k_repack_entries in + // ggml_provider.cpp, which this table mirrors label-for-label so the + // registered RepackKey (used by bench_repack.cpp for act_type/base_type) + // matches exactly. gemv/gemm repack candidates for every weight family + // GGML registers a repack entry for (Q4_0, Q8_0, IQ4_NL, MXFP4, Q4_K, + // Q5_K, Q6_K, Q2_K) are registered below. + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 4, 4, "q4_0_4x4_q8_0"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_0_4x4); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 8, 4, "q4_0_4x8_q8_0"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_0_4x8); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 8, 8, "q4_0_8x8_q8_0"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_0_4x8); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q4_K, GGML_TYPE_Q8_K, 8, 4, "q4_K_8x4_q8_K"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_k_4x4); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q4_K, GGML_TYPE_Q8_K, 8, 8, "q4_K_8x8_q8_K"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_k_4x8); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q5_K, GGML_TYPE_Q8_K, 8, 4, "q5_K_8x4_q8_K"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_k_4x4); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q5_K, GGML_TYPE_Q8_K, 8, 8, "q5_K_8x8_q8_K"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_k_4x8); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q6_K, GGML_TYPE_Q8_K, 8, 4, "q6_K_8x4_q8_K"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_k_4x4); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q6_K, GGML_TYPE_Q8_K, 8, 8, "q6_K_8x8_q8_K"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_k_4x8); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q2_K, GGML_TYPE_Q8_K, 8, 8, "q2_K_8x8_q8_K"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_k_4x8); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_0, 4, 4, "iq4_nl_4x4_q8_0"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_0_4x4); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_0, 8, 8, "iq4_nl_8x8_q8_0"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_0_4x8); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_MXFP4, GGML_TYPE_Q8_0, 4, 4, "mxfp4_4x4_q8_0"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_0_4x4); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_MXFP4, GGML_TYPE_Q8_0, 8, 8, "mxfp4_8x8_q8_0"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_0_4x8); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, 4, 4, "q8_0_4x4_q8_0"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_0_4x4); + registries.repack_quantize_mat.register_candidate({GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, 8, 4, "q8_0_4x8_q8_0"}, + "halide", ggml_quants_halide_repack_quantize_mat_q8_0_4x8); + + // gemv/gemm: same RepackKeys as the quantize_mat table above (base type, + // activation type, blocklen, ncols_interleaved, label must match + // label-for-label so bench_repack.cpp's per-key lookups line up). + registries.repack_gemv.register_candidate({GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 4, 4, "q4_0_4x4_q8_0"}, "halide", + ggml_quants_halide_repack_gemv_q4_0_4x4_q8_0); + registries.repack_gemm.register_candidate({GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 4, 4, "q4_0_4x4_q8_0"}, "halide", + ggml_quants_halide_repack_gemm_q4_0_4x4_q8_0); + registries.repack_gemv.register_candidate({GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 8, 4, "q4_0_4x8_q8_0"}, "halide", + ggml_quants_halide_repack_gemv_q4_0_4x8_q8_0); + registries.repack_gemm.register_candidate({GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 8, 4, "q4_0_4x8_q8_0"}, "halide", + ggml_quants_halide_repack_gemm_q4_0_4x8_q8_0); + registries.repack_gemv.register_candidate({GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 8, 8, "q4_0_8x8_q8_0"}, "halide", + ggml_quants_halide_repack_gemv_q4_0_8x8_q8_0); + registries.repack_gemm.register_candidate({GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, 8, 8, "q4_0_8x8_q8_0"}, "halide", + ggml_quants_halide_repack_gemm_q4_0_8x8_q8_0); + registries.repack_gemv.register_candidate({GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, 4, 4, "q8_0_4x4_q8_0"}, "halide", + ggml_quants_halide_repack_gemv_q8_0_4x4_q8_0); + registries.repack_gemm.register_candidate({GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, 4, 4, "q8_0_4x4_q8_0"}, "halide", + ggml_quants_halide_repack_gemm_q8_0_4x4_q8_0); + registries.repack_gemv.register_candidate({GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, 8, 4, "q8_0_4x8_q8_0"}, "halide", + ggml_quants_halide_repack_gemv_q8_0_4x8_q8_0); + registries.repack_gemm.register_candidate({GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, 8, 4, "q8_0_4x8_q8_0"}, "halide", + ggml_quants_halide_repack_gemm_q8_0_4x8_q8_0); + registries.repack_gemv.register_candidate({GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_0, 4, 4, "iq4_nl_4x4_q8_0"}, "halide", + ggml_quants_halide_repack_gemv_iq4_nl_4x4_q8_0); + registries.repack_gemm.register_candidate({GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_0, 4, 4, "iq4_nl_4x4_q8_0"}, "halide", + ggml_quants_halide_repack_gemm_iq4_nl_4x4_q8_0); + registries.repack_gemv.register_candidate({GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_0, 8, 8, "iq4_nl_8x8_q8_0"}, "halide", + ggml_quants_halide_repack_gemv_iq4_nl_8x8_q8_0); + registries.repack_gemm.register_candidate({GGML_TYPE_IQ4_NL, GGML_TYPE_Q8_0, 8, 8, "iq4_nl_8x8_q8_0"}, "halide", + ggml_quants_halide_repack_gemm_iq4_nl_8x8_q8_0); + registries.repack_gemv.register_candidate({GGML_TYPE_MXFP4, GGML_TYPE_Q8_0, 4, 4, "mxfp4_4x4_q8_0"}, "halide", + ggml_quants_halide_repack_gemv_mxfp4_4x4_q8_0); + registries.repack_gemm.register_candidate({GGML_TYPE_MXFP4, GGML_TYPE_Q8_0, 4, 4, "mxfp4_4x4_q8_0"}, "halide", + ggml_quants_halide_repack_gemm_mxfp4_4x4_q8_0); + registries.repack_gemv.register_candidate({GGML_TYPE_MXFP4, GGML_TYPE_Q8_0, 8, 8, "mxfp4_8x8_q8_0"}, "halide", + ggml_quants_halide_repack_gemv_mxfp4_8x8_q8_0); + registries.repack_gemm.register_candidate({GGML_TYPE_MXFP4, GGML_TYPE_Q8_0, 8, 8, "mxfp4_8x8_q8_0"}, "halide", + ggml_quants_halide_repack_gemm_mxfp4_8x8_q8_0); + registries.repack_gemv.register_candidate({GGML_TYPE_Q4_K, GGML_TYPE_Q8_K, 8, 4, "q4_K_8x4_q8_K"}, "halide", + ggml_quants_halide_repack_gemv_q4_k_8x4_q8_k); + registries.repack_gemm.register_candidate({GGML_TYPE_Q4_K, GGML_TYPE_Q8_K, 8, 4, "q4_K_8x4_q8_K"}, "halide", + ggml_quants_halide_repack_gemm_q4_k_8x4_q8_k); + registries.repack_gemv.register_candidate({GGML_TYPE_Q4_K, GGML_TYPE_Q8_K, 8, 8, "q4_K_8x8_q8_K"}, "halide", + ggml_quants_halide_repack_gemv_q4_k_8x8_q8_k); + registries.repack_gemm.register_candidate({GGML_TYPE_Q4_K, GGML_TYPE_Q8_K, 8, 8, "q4_K_8x8_q8_K"}, "halide", + ggml_quants_halide_repack_gemm_q4_k_8x8_q8_k); + registries.repack_gemv.register_candidate({GGML_TYPE_Q5_K, GGML_TYPE_Q8_K, 8, 4, "q5_K_8x4_q8_K"}, "halide", + ggml_quants_halide_repack_gemv_q5_k_8x4_q8_k); + registries.repack_gemm.register_candidate({GGML_TYPE_Q5_K, GGML_TYPE_Q8_K, 8, 4, "q5_K_8x4_q8_K"}, "halide", + ggml_quants_halide_repack_gemm_q5_k_8x4_q8_k); + registries.repack_gemv.register_candidate({GGML_TYPE_Q5_K, GGML_TYPE_Q8_K, 8, 8, "q5_K_8x8_q8_K"}, "halide", + ggml_quants_halide_repack_gemv_q5_k_8x8_q8_k); + registries.repack_gemm.register_candidate({GGML_TYPE_Q5_K, GGML_TYPE_Q8_K, 8, 8, "q5_K_8x8_q8_K"}, "halide", + ggml_quants_halide_repack_gemm_q5_k_8x8_q8_k); + registries.repack_gemv.register_candidate({GGML_TYPE_Q6_K, GGML_TYPE_Q8_K, 8, 4, "q6_K_8x4_q8_K"}, "halide", + ggml_quants_halide_repack_gemv_q6_k_8x4_q8_k); + registries.repack_gemm.register_candidate({GGML_TYPE_Q6_K, GGML_TYPE_Q8_K, 8, 4, "q6_K_8x4_q8_K"}, "halide", + ggml_quants_halide_repack_gemm_q6_k_8x4_q8_k); + registries.repack_gemv.register_candidate({GGML_TYPE_Q6_K, GGML_TYPE_Q8_K, 8, 8, "q6_K_8x8_q8_K"}, "halide", + ggml_quants_halide_repack_gemv_q6_k_8x8_q8_k); + registries.repack_gemm.register_candidate({GGML_TYPE_Q6_K, GGML_TYPE_Q8_K, 8, 8, "q6_K_8x8_q8_K"}, "halide", + ggml_quants_halide_repack_gemm_q6_k_8x8_q8_k); + registries.repack_gemv.register_candidate({GGML_TYPE_Q2_K, GGML_TYPE_Q8_K, 8, 8, "q2_K_8x8_q8_K"}, "halide", + ggml_quants_halide_repack_gemv_q2_k_8x8_q8_k); + registries.repack_gemm.register_candidate({GGML_TYPE_Q2_K, GGML_TYPE_Q8_K, 8, 8, "q2_K_8x8_q8_K"}, "halide", + ggml_quants_halide_repack_gemm_q2_k_8x8_q8_k); +} diff --git a/apps/ggml/providers/halide_provider.h b/apps/ggml/providers/halide_provider.h new file mode 100644 index 000000000000..02b80d345536 --- /dev/null +++ b/apps/ggml/providers/halide_provider.h @@ -0,0 +1,8 @@ +#pragma once + +#include "kernel_registry.h" + +// Registers the from-scratch Halide reimplementation of GGML's Q4_0 +// quantize/dequantize kernels (see ../halide/) as candidates against GGML's +// own reference, which register_ggml_provider() registers first. +void register_halide_provider(KernelRegistries ®istries); diff --git a/apps/ggml/src/bench_dequantize.cpp b/apps/ggml/src/bench_dequantize.cpp new file mode 100644 index 000000000000..260645b80656 --- /dev/null +++ b/apps/ggml/src/bench_dequantize.cpp @@ -0,0 +1,67 @@ +#include "benchmarks.h" + +#include + +#include "compare.h" +#include "data_gen.h" +#include "timing.h" + +namespace { +constexpr int64_t kTargetElements = 4096; +} // namespace + +BenchReport run_dequantize_benchmarks(const KernelRegistries ®istries) { + BenchReport report{"dequantize_row", "GB/s", {}}; + print_report_header(report.title, report.throughput_unit); + + for (int t = 0; t < GGML_TYPE_COUNT; ++t) { + const ggml_type type = static_cast(t); + const Impl *ref = registries.dequantize.reference(type); + const Impl *qref = registries.quantize.reference(type); + if (!ref || !qref) { + continue; + } + ggml_quantize_init(type); // one-time, cheap after the first call for this type; see ggml_provider.cpp + + const int64_t blck = ggml_blck_size(type); + const int64_t n = ((kTargetElements + blck - 1) / blck) * blck; + + AlignedBuffer src(n * sizeof(float)); + generate_synthetic_data(src.as(), n); + + AlignedBuffer quantized(ggml_row_size(type, n)); + qref->fn(src.as(), quantized.data(), n); + + AlignedBuffer ref_out(n * sizeof(float)); + ref->fn(quantized.data(), ref_out.as(), n); + const TimingResult ref_time = time_calls([&] { ref->fn(quantized.data(), ref_out.as(), n); }); + + BenchRow row; + row.label = ggml_type_name(type); + row.ref_name = ref->name; + row.ref_ns = ref_time.median_ns; + row.ref_throughput = bytes_per_sec(n * sizeof(float), ref_time.median_ns) / 1e9; + + for (const auto &cand : registries.dequantize.candidates(type)) { + BenchCandidate bc; + bc.name = cand.name; + bc.identical = (cand.fn == ref->fn); + if (!bc.identical) { + AlignedBuffer cand_out(n * sizeof(float)); + cand.fn(quantized.data(), cand_out.as(), n); + bc.correct = floats_match(ref_out.as(), cand_out.as(), n); + + const TimingResult cand_time = time_calls([&] { cand.fn(quantized.data(), cand_out.as(), n); }); + bc.ns = cand_time.median_ns; + bc.throughput = bytes_per_sec(n * sizeof(float), cand_time.median_ns) / 1e9; + bc.speedup = ref_time.median_ns / cand_time.median_ns; + } + row.candidates.push_back(bc); + } + + print_row(row, report.throughput_unit); + report.rows.push_back(std::move(row)); + } + + return report; +} diff --git a/apps/ggml/src/bench_quantize.cpp b/apps/ggml/src/bench_quantize.cpp new file mode 100644 index 000000000000..56d60dcd7f71 --- /dev/null +++ b/apps/ggml/src/bench_quantize.cpp @@ -0,0 +1,65 @@ +#include "benchmarks.h" + +#include + +#include + +#include "data_gen.h" +#include "timing.h" + +namespace { +constexpr int64_t kTargetElements = 4096; +} + +BenchReport run_quantize_benchmarks(const KernelRegistries ®istries) { + BenchReport report{"quantize_row", "GB/s", {}}; + print_report_header(report.title, report.throughput_unit); + + for (int t = 0; t < GGML_TYPE_COUNT; ++t) { + const ggml_type type = static_cast(t); + const Impl *ref = registries.quantize.reference(type); + if (!ref) { + continue; + } + ggml_quantize_init(type); // one-time, cheap after the first call for this type; see ggml_provider.cpp + + const int64_t blck = ggml_blck_size(type); + const int64_t n = ((kTargetElements + blck - 1) / blck) * blck; + const size_t out_bytes = ggml_row_size(type, n); + + AlignedBuffer src(n * sizeof(float)); + generate_synthetic_data(src.as(), n); + + AlignedBuffer ref_out(out_bytes); + ref->fn(src.as(), ref_out.data(), n); + const TimingResult ref_time = time_calls([&] { ref->fn(src.as(), ref_out.data(), n); }); + + BenchRow row; + row.label = ggml_type_name(type); + row.ref_name = ref->name; + row.ref_ns = ref_time.median_ns; + row.ref_throughput = bytes_per_sec(n * sizeof(float), ref_time.median_ns) / 1e9; + + for (const auto &cand : registries.quantize.candidates(type)) { + BenchCandidate bc; + bc.name = cand.name; + bc.identical = (cand.fn == ref->fn); + if (!bc.identical) { + AlignedBuffer cand_out(out_bytes); + cand.fn(src.as(), cand_out.data(), n); + bc.correct = (std::memcmp(ref_out.data(), cand_out.data(), out_bytes) == 0); + + const TimingResult cand_time = time_calls([&] { cand.fn(src.as(), cand_out.data(), n); }); + bc.ns = cand_time.median_ns; + bc.throughput = bytes_per_sec(n * sizeof(float), cand_time.median_ns) / 1e9; + bc.speedup = ref_time.median_ns / cand_time.median_ns; + } + row.candidates.push_back(bc); + } + + print_row(row, report.throughput_unit); + report.rows.push_back(std::move(row)); + } + + return report; +} diff --git a/apps/ggml/src/bench_repack.cpp b/apps/ggml/src/bench_repack.cpp new file mode 100644 index 000000000000..18ca3b675222 --- /dev/null +++ b/apps/ggml/src/bench_repack.cpp @@ -0,0 +1,269 @@ +#include "benchmarks.h" + +#include +#include + +#include +#include +#include +#include + +#include "compare.h" +#include "data_gen.h" +#include "ggml_internal_abi.h" // ggml_backend_cpu_repack_buffer_type() +#include "timing.h" + +namespace { + +// K (reduction dim): divisible by every block size in play (32 and 256). +constexpr int64_t kK = 4096; +// Output columns: divisible by every NB_COLS in play (4 and 8). +constexpr int kNC = 32; +// Activation rows for the gemm (batched) path; must be a multiple of 4 (the +// row-group size ggml_quantize_mat_* always packs), and > 3 so production +// code would pick gemm over gemv for this many rows (see repack.cpp's +// forward_mul_mat_one_chunk: "if there are more than three rows in src1, +// use gemm; otherwise, use gemv"). +constexpr int kGemmRows = 32; + +// Builds a packed weight buffer for `base_type`, shape [kK, kNC], by +// quantizing a row-major staging buffer through the reference quantizer and +// then letting the repack buffer type's set_tensor callback do the actual +// interleaving (src/ggml-cpu/repack.cpp:4733). This avoids reimplementing +// the private block interleave layout by hand -- everything here is +// public API plus the one declared-ourselves accessor for the buffer type. +struct PackedWeight { + ggml_context *ctx = nullptr; + ggml_backend_buffer_t buffer = nullptr; + ggml_tensor *tensor = nullptr; + + PackedWeight() = default; + PackedWeight(const PackedWeight &) = delete; + PackedWeight &operator=(const PackedWeight &) = delete; + PackedWeight(PackedWeight &&other) noexcept + : ctx(other.ctx), buffer(other.buffer), tensor(other.tensor) { + other.ctx = nullptr; + other.buffer = nullptr; + other.tensor = nullptr; + } + PackedWeight &operator=(PackedWeight &&other) noexcept { + if (this != &other) { + if (buffer) ggml_backend_buffer_free(buffer); + if (ctx) ggml_free(ctx); + ctx = other.ctx; + buffer = other.buffer; + tensor = other.tensor; + other.ctx = nullptr; + other.buffer = nullptr; + other.tensor = nullptr; + } + return *this; + } + ~PackedWeight() { + if (buffer) ggml_backend_buffer_free(buffer); + if (ctx) ggml_free(ctx); + } + + // ggml_repack_get_optimal_repack_type() (src/ggml-cpu/repack.cpp) picks + // ONE interleave layout per (type, CPU features, column count) -- it is + // GGML's own "best kernel for this hardware" heuristic, not something + // this benchmark can steer towards a specific registered combo. When it + // finds none (e.g. Q2_K has no ARM branch at all, only AVX512/RISC-V -- + // see that function), the buffer type's init_tensor callback leaves + // tensor->extra null, and calling set_tensor on it would dereference a + // null tensor_traits pointer. There is no supported way to force a + // different combo through this public mechanism, so such combos are + // skipped rather than benchmarked with fabricated data. + bool supported_on_this_cpu() const { + return tensor && tensor->extra != nullptr; + } +}; + +PackedWeight build_packed_weight(ggml_type base_type, const Impl &row_quant, int64_t k, int nc) { + PackedWeight pw; + ggml_init_params params{/*.mem_size=*/ggml_tensor_overhead() + 256, /*.mem_buffer=*/nullptr, + /*.no_alloc=*/true}; + pw.ctx = ggml_init(params); + pw.tensor = ggml_new_tensor_2d(pw.ctx, base_type, k, nc); + pw.buffer = ggml_backend_alloc_ctx_tensors_from_buft(pw.ctx, ggml_backend_cpu_repack_buffer_type()); + if (!pw.supported_on_this_cpu()) { + return pw; + } + + AlignedBuffer staging(ggml_row_size(base_type, k) * nc); + AlignedBuffer col(k * sizeof(float)); + const size_t row_bytes = ggml_row_size(base_type, k); + for (int c = 0; c < nc; ++c) { + generate_synthetic_data(col.as(), k, static_cast(c)); + row_quant.fn(col.as(), staging.as() + c * row_bytes, k); + } + ggml_backend_tensor_set(pw.tensor, staging.data(), 0, staging.size()); + return pw; +} + +} // namespace + +std::vector run_repack_benchmarks(const KernelRegistries ®istries) { + BenchReport quant_mat_report{"repack_quantize_mat", "GB/s", {}}; + BenchReport gemv_report{"repack_gemv", "GFLOP/s", {}}; + BenchReport gemm_report{"repack_gemm", "GFLOP/s", {}}; + + print_report_header(quant_mat_report.title, quant_mat_report.throughput_unit); + for (const RepackKey &key : registries.repack_quantize_mat.keys()) { + const Impl *qm_ref = registries.repack_quantize_mat.reference(key.label); + if (!qm_ref) { + continue; + } + const int64_t blck = ggml_blck_size(key.act_type); + const int64_t k = ((kK + blck - 1) / blck) * blck; + + // ggml_quantize_mat_* always consumes exactly 4 rows (see + // src/ggml-cpu/repack.cpp: ggml_quantize_mat_t<...> asserts nrow==4), + // regardless of the interleave geometry in the name. + AlignedBuffer src(4 * k * sizeof(float)); + generate_synthetic_data(src.as(), 4 * k); + const size_t out_bytes = 4 * ggml_row_size(key.act_type, k); + + AlignedBuffer ref_out(out_bytes); + qm_ref->fn(src.as(), ref_out.data(), k); + const TimingResult ref_time = time_calls([&] { qm_ref->fn(src.as(), ref_out.data(), k); }); + + BenchRow row; + row.label = key.label; + row.ref_name = qm_ref->name; + row.ref_ns = ref_time.median_ns; + row.ref_throughput = bytes_per_sec(4 * k * sizeof(float), ref_time.median_ns) / 1e9; + + for (const auto &cand : registries.repack_quantize_mat.candidates(key.label)) { + BenchCandidate bc; + bc.name = cand.name; + bc.identical = (cand.fn == qm_ref->fn); + if (!bc.identical) { + AlignedBuffer cand_out(out_bytes); + cand.fn(src.as(), cand_out.data(), k); + bc.correct = (std::memcmp(ref_out.data(), cand_out.data(), out_bytes) == 0); + + const TimingResult cand_time = time_calls([&] { cand.fn(src.as(), cand_out.data(), k); }); + bc.ns = cand_time.median_ns; + bc.throughput = bytes_per_sec(4 * k * sizeof(float), cand_time.median_ns) / 1e9; + bc.speedup = ref_time.median_ns / cand_time.median_ns; + } + row.candidates.push_back(bc); + } + print_row(row, quant_mat_report.throughput_unit); + quant_mat_report.rows.push_back(std::move(row)); + } + + print_report_header(gemv_report.title, gemv_report.throughput_unit); + for (const RepackKey &key : registries.repack_gemv.keys()) { + const Impl *gemv_ref = registries.repack_gemv.reference(key.label); + const Impl *w_quant = registries.quantize.reference(key.base_type); + const Impl *a_quant = registries.quantize.reference(key.act_type); + if (!gemv_ref || !w_quant || !a_quant) { + continue; + } + const int64_t blck = ggml_blck_size(key.base_type); + const int64_t k = ((kK + blck - 1) / blck) * blck; + + PackedWeight weight = build_packed_weight(key.base_type, *w_quant, k, kNC); + if (!weight.supported_on_this_cpu()) { + continue; // no repack kernel for this (type, CPU) combo -- see PackedWeight::supported_on_this_cpu + } + + AlignedBuffer a_src(k * sizeof(float)); + generate_synthetic_data(a_src.as(), k, 3.0f); + AlignedBuffer vy(ggml_row_size(key.act_type, k)); + a_quant->fn(a_src.as(), vy.data(), k); + + AlignedBuffer ref_out(kNC * sizeof(float)); + gemv_ref->fn(k, ref_out.as(), kNC, weight.tensor->data, vy.data(), 1, kNC); + const TimingResult ref_time = + time_calls([&] { gemv_ref->fn(k, ref_out.as(), kNC, weight.tensor->data, vy.data(), 1, kNC); }); + const double ref_flops = 2.0 * k * kNC; + + BenchRow row; + row.label = key.label; + row.ref_name = gemv_ref->name; + row.ref_ns = ref_time.median_ns; + row.ref_throughput = gflops(ref_flops, ref_time.median_ns); + + for (const auto &cand : registries.repack_gemv.candidates(key.label)) { + BenchCandidate bc; + bc.name = cand.name; + bc.identical = (cand.fn == gemv_ref->fn); + if (!bc.identical) { + AlignedBuffer cand_out(kNC * sizeof(float)); + cand.fn(k, cand_out.as(), kNC, weight.tensor->data, vy.data(), 1, kNC); + bc.correct = floats_match(ref_out.as(), cand_out.as(), kNC); + + const TimingResult cand_time = + time_calls([&] { cand.fn(k, cand_out.as(), kNC, weight.tensor->data, vy.data(), 1, kNC); }); + bc.ns = cand_time.median_ns; + bc.throughput = gflops(ref_flops, cand_time.median_ns); + bc.speedup = ref_time.median_ns / cand_time.median_ns; + } + row.candidates.push_back(bc); + } + print_row(row, gemv_report.throughput_unit); + gemv_report.rows.push_back(std::move(row)); + } + + print_report_header(gemm_report.title, gemm_report.throughput_unit); + for (const RepackKey &key : registries.repack_gemm.keys()) { + const Impl *gemm_ref = registries.repack_gemm.reference(key.label); + const Impl *w_quant = registries.quantize.reference(key.base_type); + const Impl *qm_ref = registries.repack_quantize_mat.reference(key.label); + if (!gemm_ref || !w_quant || !qm_ref) { + continue; + } + const int64_t blck = ggml_blck_size(key.base_type); + const int64_t k = ((kK + blck - 1) / blck) * blck; + + PackedWeight weight = build_packed_weight(key.base_type, *w_quant, k, kNC); + if (!weight.supported_on_this_cpu()) { + continue; // no repack kernel for this (type, CPU) combo -- see PackedWeight::supported_on_this_cpu + } + + AlignedBuffer a_src(kGemmRows * k * sizeof(float)); + generate_synthetic_data(a_src.as(), kGemmRows * k, 5.0f); + const size_t group_bytes = 4 * ggml_row_size(key.act_type, k); + AlignedBuffer vy(group_bytes * (kGemmRows / 4)); + for (int g = 0; g < kGemmRows / 4; ++g) { + qm_ref->fn(a_src.as() + g * 4 * k, vy.as() + g * group_bytes, k); + } + + AlignedBuffer ref_out(kGemmRows * kNC * sizeof(float)); + gemm_ref->fn(k, ref_out.as(), kNC, weight.tensor->data, vy.data(), kGemmRows, kNC); + const TimingResult ref_time = time_calls( + [&] { gemm_ref->fn(k, ref_out.as(), kNC, weight.tensor->data, vy.data(), kGemmRows, kNC); }); + const double ref_flops = 2.0 * k * kNC * kGemmRows; + + BenchRow row; + row.label = key.label; + row.ref_name = gemm_ref->name; + row.ref_ns = ref_time.median_ns; + row.ref_throughput = gflops(ref_flops, ref_time.median_ns); + + for (const auto &cand : registries.repack_gemm.candidates(key.label)) { + BenchCandidate bc; + bc.name = cand.name; + bc.identical = (cand.fn == gemm_ref->fn); + if (!bc.identical) { + AlignedBuffer cand_out(kGemmRows * kNC * sizeof(float)); + cand.fn(k, cand_out.as(), kNC, weight.tensor->data, vy.data(), kGemmRows, kNC); + bc.correct = floats_match(ref_out.as(), cand_out.as(), kGemmRows * kNC); + + const TimingResult cand_time = time_calls( + [&] { cand.fn(k, cand_out.as(), kNC, weight.tensor->data, vy.data(), kGemmRows, kNC); }); + bc.ns = cand_time.median_ns; + bc.throughput = gflops(ref_flops, cand_time.median_ns); + bc.speedup = ref_time.median_ns / cand_time.median_ns; + } + row.candidates.push_back(bc); + } + print_row(row, gemm_report.throughput_unit); + gemm_report.rows.push_back(std::move(row)); + } + + return {quant_mat_report, gemv_report, gemm_report}; +} diff --git a/apps/ggml/src/bench_vecdot.cpp b/apps/ggml/src/bench_vecdot.cpp new file mode 100644 index 000000000000..c2e522794d28 --- /dev/null +++ b/apps/ggml/src/bench_vecdot.cpp @@ -0,0 +1,116 @@ +#include "benchmarks.h" + +#include +#include + +#include "compare.h" +#include "data_gen.h" +#include "timing.h" + +#include +#include + +namespace { +// Divisible by every quant block size in play (32 for the q4/q5/q8 family, 256 for the k-quants). +// Dev iteration aid: KERNEL_BENCH_N overrides it (must stay a multiple of 256), +// which is how per-call overhead is separated from per-block cost. +int64_t elements() { + const char *n = std::getenv("KERNEL_BENCH_N"); + return (n && *n) ? std::atoll(n) : 4096; +} +const int64_t kElements = elements(); + +// Dev iteration aid: KERNEL_BENCH_FILTER=q4_0,q8_0 restricts the sweep to +// matching type names (substring match, comma-separated). Empty => all. +bool passes_filter(const char *name) { + const char *f = std::getenv("KERNEL_BENCH_FILTER"); + if (!f || !*f) { + return true; + } + std::string filt(f), item; + size_t pos = 0; + while (pos <= filt.size()) { + size_t comma = filt.find(',', pos); + if (comma == std::string::npos) comma = filt.size(); + item = filt.substr(pos, comma - pos); + if (!item.empty() && std::strstr(name, item.c_str())) { + return true; + } + pos = comma + 1; + } + return false; +} +} // namespace + +BenchReport run_vecdot_benchmarks(const KernelRegistries ®istries) { + BenchReport report{"vec_dot", "GB/s", {}}; + print_report_header(report.title, report.throughput_unit); + + for (int t = 0; t < GGML_TYPE_COUNT; ++t) { + const ggml_type type = static_cast(t); + const Impl *ref = registries.vec_dot.reference(type); + if (!ref) { + continue; + } + if (!passes_filter(ggml_type_name(type))) { + continue; + } + + const ggml_type_traits_cpu *tc = ggml_get_type_traits_cpu(type); + if (!tc) { + continue; + } + const ggml_type act_type = tc->vec_dot_type; + + const Impl *x_quant = registries.quantize.reference(type); + const Impl *y_quant = registries.quantize.reference(act_type); + if (!x_quant || !y_quant) { + continue; // shouldn't happen for any type reachable through the CPU backend + } + ggml_quantize_init(type); // one-time, cheap after the first call for this type; see ggml_provider.cpp + + AlignedBuffer x_src(kElements * sizeof(float)); + AlignedBuffer y_src(kElements * sizeof(float)); + generate_synthetic_data(x_src.as(), kElements, 0.0f); + generate_synthetic_data(y_src.as(), kElements, 7.0f); // different phase so x != y + + AlignedBuffer vx(ggml_row_size(type, kElements)); + AlignedBuffer vy(ggml_row_size(act_type, kElements)); + x_quant->fn(x_src.as(), vx.data(), kElements); + y_quant->fn(y_src.as(), vy.data(), kElements); + + float ref_result = 0.0f; + ref->fn(kElements, &ref_result, 0, vx.data(), 0, vy.data(), 0, 1); + const TimingResult ref_time = + time_calls([&] { ref->fn(kElements, &ref_result, 0, vx.data(), 0, vy.data(), 0, 1); }); + + BenchRow row; + row.label = ggml_type_name(type); + row.ref_name = ref->name; + row.ref_ns = ref_time.min_ns; + row.ref_throughput = bytes_per_sec(vx.size() + vy.size(), ref_time.min_ns) / 1e9; + + for (const auto &cand : registries.vec_dot.candidates(type)) { + BenchCandidate bc; + bc.name = cand.name; + bc.identical = (cand.fn == ref->fn); + if (!bc.identical) { + float cand_result = 0.0f; + cand.fn(kElements, &cand_result, 0, vx.data(), 0, vy.data(), 0, 1); + bc.correct = floats_match(&ref_result, &cand_result, 1); + + const TimingResult cand_time = + time_calls([&] { cand.fn(kElements, &cand_result, 0, vx.data(), 0, vy.data(), 0, 1); }); + bc.ns = cand_time.min_ns; + bc.throughput = bytes_per_sec(vx.size() + vy.size(), cand_time.min_ns) / 1e9; + bc.speedup = ref_time.min_ns / cand_time.min_ns; + } + row.candidates.push_back(bc); + } + + print_row(row, report.throughput_unit); + report.rows.push_back(std::move(row)); + } + + return report; +} diff --git a/apps/ggml/src/benchmarks.h b/apps/ggml/src/benchmarks.h new file mode 100644 index 000000000000..75029adbc44d --- /dev/null +++ b/apps/ggml/src/benchmarks.h @@ -0,0 +1,13 @@ +#pragma once + +#include + +#include "kernel_registry.h" +#include "report.h" + +BenchReport run_quantize_benchmarks(const KernelRegistries ®istries); +BenchReport run_dequantize_benchmarks(const KernelRegistries ®istries); +BenchReport run_vecdot_benchmarks(const KernelRegistries ®istries); + +// One report each for quantize_mat, gemv, gemm. +std::vector run_repack_benchmarks(const KernelRegistries ®istries); diff --git a/apps/ggml/src/compare.h b/apps/ggml/src/compare.h new file mode 100644 index 000000000000..b056c3c0d8fe --- /dev/null +++ b/apps/ggml/src/compare.h @@ -0,0 +1,20 @@ +#pragma once + +#include +#include +#include + +// Relative-error comparison for floating point outputs (dequantize, vec_dot, +// gemv/gemm results). Quantize/quantize_mat outputs are compared with an +// exact memcmp instead (see bench_quantize.cpp) since those algorithms are +// specified to be bit-identical across implementations. +inline bool floats_match(const float *a, const float *b, int64_t n, float rel_tol = 1e-2f) { + for (int64_t i = 0; i < n; ++i) { + const float diff = std::fabs(a[i] - b[i]); + const float scale = std::max(std::fabs(a[i]), 1e-6f); + if (diff / scale > rel_tol) { + return false; + } + } + return true; +} diff --git a/apps/ggml/src/data_gen.h b/apps/ggml/src/data_gen.h new file mode 100644 index 000000000000..22075c369abf --- /dev/null +++ b/apps/ggml/src/data_gen.h @@ -0,0 +1,60 @@ +#pragma once + +// Deterministic synthetic data + aligned buffers, following the conventions +// of tests/test-quantize-perf.cpp (same generator, same rationale: a fixed +// seedless formula so every implementation under comparison sees byte-identical +// input without carrying a PRNG dependency). + +#include +#include +#include +#include +#include + +inline void generate_synthetic_data(float *dst, size_t n, float offset = 0.0f) { + for (size_t i = 0; i < n; ++i) { + dst[i] = 0.1f + 2.0f * cosf(static_cast(i) + offset); + } +} + +// 64-byte aligned heap buffer (covers every SIMD width in use: SSE/AVX/AVX-512/NEON/SVE). +class AlignedBuffer { +public: + explicit AlignedBuffer(size_t bytes) : size_(bytes) { + constexpr size_t alignment = 64; + size_t padded = (bytes + alignment - 1) / alignment * alignment; + if (padded == 0) { + padded = alignment; + } + ptr_ = nullptr; + posix_memalign(&ptr_, alignment, padded); + } + ~AlignedBuffer() { + std::free(ptr_); + } + + AlignedBuffer(const AlignedBuffer &) = delete; + AlignedBuffer &operator=(const AlignedBuffer &) = delete; + + void *data() { + return ptr_; + } + const void *data() const { + return ptr_; + } + template + T *as() { + return static_cast(ptr_); + } + template + const T *as() const { + return static_cast(ptr_); + } + size_t size() const { + return size_; + } + +private: + void *ptr_; + size_t size_; +}; diff --git a/apps/ggml/src/main.cpp b/apps/ggml/src/main.cpp new file mode 100644 index 000000000000..24cb441426f8 --- /dev/null +++ b/apps/ggml/src/main.cpp @@ -0,0 +1,122 @@ +#include +#include +#include + +#include + +#include "benchmarks.h" +#include "ggml_provider.h" +#include "halide_provider.h" +#include "kernel_registry.h" + +namespace { + +// The repack buffer type logs a GGML_LOG_DEBUG line on every repack (see +// src/ggml-cpu/repack.cpp:4733) -- benign, but this benchmark triggers many +// of them (one per gemv/gemm sample weight built), so drop DEBUG/INFO noise +// and keep only warnings/errors. +void quiet_log_callback(ggml_log_level level, const char *text, void *) { + if (level >= GGML_LOG_LEVEL_WARN) { + std::fputs(text, stderr); + } +} + +void print_ggml_version() { +#ifdef KERNEL_BENCH_GGML_VERSION + std::printf("GGML version: %s\n", KERNEL_BENCH_GGML_VERSION); +#else + std::printf("GGML version: unknown (GGML_VERSION not set by ggml-config.cmake)\n"); +#endif +} + +void print_cpu_features() { + std::printf("CPU features:"); +#if defined(__x86_64__) || defined(__i386__) || defined(_M_IX86) || defined(_M_X64) + if (ggml_cpu_has_avx()) std::printf(" avx"); + if (ggml_cpu_has_avx2()) std::printf(" avx2"); + if (ggml_cpu_has_avx512()) std::printf(" avx512"); + if (ggml_cpu_has_avx512_vnni()) std::printf(" avx512_vnni"); + if (ggml_cpu_has_fma()) std::printf(" fma"); + if (ggml_cpu_has_f16c()) std::printf(" f16c"); + if (ggml_cpu_has_amx_int8()) std::printf(" amx_int8"); +#endif +#if defined(__aarch64__) || defined(__arm__) || defined(_M_ARM) || defined(_M_ARM64) + if (ggml_cpu_has_neon()) std::printf(" neon"); + if (ggml_cpu_has_dotprod()) std::printf(" dotprod"); + if (ggml_cpu_has_matmul_int8()) std::printf(" matmul_int8"); + if (ggml_cpu_has_fp16_va()) std::printf(" fp16_va"); + if (ggml_cpu_has_sve()) std::printf(" sve(%d bytes)", ggml_cpu_get_sve_cnt()); + if (ggml_cpu_has_sme()) std::printf(" sme"); // codespell:ignore sme +#endif + std::printf("\n"); +} + +void usage(const char *argv0) { + std::printf("usage: %s [--quantize] [--dequantize] [--vecdot] [--repack] [--all] [--csv FILE]\n", argv0); +} + +} // namespace + +int main(int argc, char **argv) { + ggml_log_set(quiet_log_callback, nullptr); + + bool do_quantize = false, do_dequantize = false, do_vecdot = false, do_repack = false; + std::string csv_path; + + for (int i = 1; i < argc; ++i) { + const std::string arg = argv[i]; + if (arg == "--quantize") do_quantize = true; + else if (arg == "--dequantize") + do_dequantize = true; + else if (arg == "--vecdot") + do_vecdot = true; + else if (arg == "--repack") + do_repack = true; + else if (arg == "--all") + do_quantize = do_dequantize = do_vecdot = do_repack = true; + else if (arg == "--csv" && i + 1 < argc) + csv_path = argv[++i]; + else if (arg == "--help" || arg == "-h") { + usage(argv[0]); + return 0; + } else { + std::fprintf(stderr, "unknown argument: %s\n", arg.c_str()); + usage(argv[0]); + return 1; + } + } + if (!do_quantize && !do_dequantize && !do_vecdot && !do_repack) { + do_quantize = do_dequantize = do_vecdot = do_repack = true; // default: --all + } + + print_ggml_version(); + print_cpu_features(); + + KernelRegistries registries; + register_ggml_provider(registries); + register_halide_provider(registries); + + // Each run_*_benchmarks() call prints its own header and streams a row + // to stdout as soon as that row is computed (see report.h/print_row) -- + // results become visible immediately rather than only after everything + // finishes. The returned reports are only needed here for --csv. + std::vector reports; + if (do_quantize) reports.push_back(run_quantize_benchmarks(registries)); + if (do_dequantize) reports.push_back(run_dequantize_benchmarks(registries)); + if (do_vecdot) reports.push_back(run_vecdot_benchmarks(registries)); + if (do_repack) { + for (auto &r : run_repack_benchmarks(registries)) { + reports.push_back(std::move(r)); + } + } + + if (!csv_path.empty()) { + std::ofstream out(csv_path); + for (const auto &report : reports) { + write_report_csv(report, out); + } + std::printf("\nwrote %s\n", csv_path.c_str()); + } + + return 0; +} diff --git a/apps/ggml/src/report.cpp b/apps/ggml/src/report.cpp new file mode 100644 index 000000000000..6527b10967ed --- /dev/null +++ b/apps/ggml/src/report.cpp @@ -0,0 +1,50 @@ +#include "report.h" + +#include +#include + +void print_report_header(const std::string &title, const std::string &throughput_unit) { + std::printf("\n=== %s (%s) ===\n", title.c_str(), throughput_unit.c_str()); + std::fflush(stdout); +} + +void print_row(const BenchRow &row, const std::string &throughput_unit) { + std::printf(" %-20s reference=%-10s %8.1f ns %8.2f %s\n", row.label.c_str(), row.ref_name.c_str(), row.ref_ns, + row.ref_throughput, throughput_unit.c_str()); + if (row.candidates.empty()) { + std::printf(" %-20s (no candidates registered yet)\n", ""); + } + for (const auto &c : row.candidates) { + if (c.identical) { + std::printf(" %-20s %-12s identical to reference\n", "", c.name.c_str()); + continue; + } + std::printf(" %-20s %-12s %8.1f ns %8.2f %s %6.2fx%s\n", "", c.name.c_str(), c.ns, c.throughput, + throughput_unit.c_str(), c.speedup, c.correct ? "" : " [MISMATCH vs reference]"); + } + std::fflush(stdout); +} + +void print_report(const BenchReport &report) { + print_report_header(report.title, report.throughput_unit); + if (report.rows.empty()) { + std::printf(" (nothing registered)\n"); + std::fflush(stdout); + return; + } + for (const auto &row : report.rows) { + print_row(row, report.throughput_unit); + } +} + +void write_report_csv(const BenchReport &report, std::ostream &out) { + out << "table,label,role,name,ns,throughput_" << report.throughput_unit << ",speedup,identical,correct\n"; + for (const auto &row : report.rows) { + out << report.title << ',' << row.label << ",reference," << row.ref_name << ',' << row.ref_ns << ',' + << row.ref_throughput << ",1.0,0,1\n"; + for (const auto &c : row.candidates) { + out << report.title << ',' << row.label << ",candidate," << c.name << ',' << c.ns << ',' << c.throughput + << ',' << c.speedup << ',' << (c.identical ? 1 : 0) << ',' << (c.correct ? 1 : 0) << '\n'; + } + } +} diff --git a/apps/ggml/src/report.h b/apps/ggml/src/report.h new file mode 100644 index 000000000000..265cab2e10eb --- /dev/null +++ b/apps/ggml/src/report.h @@ -0,0 +1,42 @@ +#pragma once + +#include +#include +#include + +struct BenchCandidate { + std::string name; + double ns = 0.0; + double throughput = 0.0; + double speedup = 0.0; // reference.ns / ns + bool identical = false; // candidate fn pointer == reference fn pointer + bool correct = true; // output matched the reference within tolerance +}; + +struct BenchRow { + std::string label; // type name, or repack key label + double ref_ns = 0.0; + double ref_throughput = 0.0; + std::string ref_name; + std::vector candidates; +}; + +struct BenchReport { + std::string title; + std::string throughput_unit; // "GB/s" or "GFLOP/s" + std::vector rows; +}; + +// Incremental printing: call print_report_header() once, then print_row() +// as each row is computed (bench_*.cpp interleaves this with the actual +// benchmarking so results stream out immediately instead of only appearing +// after the whole category finishes). Both flush stdout so the stream is +// visible immediately even when redirected/piped, not just on a tty. +void print_report_header(const std::string &title, const std::string &throughput_unit); +void print_row(const BenchRow &row, const std::string &throughput_unit); + +// Convenience wrapper for a fully-built report (used for the "nothing +// registered" case, and anywhere the whole report is already in hand). +void print_report(const BenchReport &report); + +void write_report_csv(const BenchReport &report, std::ostream &out); diff --git a/apps/ggml/src/timing.h b/apps/ggml/src/timing.h new file mode 100644 index 000000000000..ad78e859c6b8 --- /dev/null +++ b/apps/ggml/src/timing.h @@ -0,0 +1,73 @@ +#pragma once + +#include +#include +#include +#include +#include + +struct TimingResult { + double min_ns = 0.0; + double median_ns = 0.0; +}; + +// Runs `fn` `warmup` times (discarded), then calibrates a batch size large +// enough that timing a whole batch back-to-back amortizes clock overhead and +// resolution (individual quantize_row/vec_dot calls on fast SIMD kernels can +// complete in a few nanoseconds -- timing them one at a time, even with a +// high-resolution clock, is dominated by noise), then times `iters` such +// batches and returns the min/median per-call latency in nanoseconds. +inline TimingResult time_calls(const std::function &fn, int warmup = 5, int iters = 20) { + using clock = std::chrono::steady_clock; + + for (int i = 0; i < warmup; ++i) { + fn(); + } + + constexpr double kMinBatchNs = 200000.0; // 0.2ms per batch + int batch = 1; + for (;;) { + const auto t0 = clock::now(); + for (int i = 0; i < batch; ++i) { + fn(); + } + const auto t1 = clock::now(); + const double batch_ns = std::chrono::duration(t1 - t0).count(); + if (batch_ns >= kMinBatchNs || batch >= (1 << 20)) { + break; + } + batch *= 4; + } + + std::vector samples_ns; + samples_ns.reserve(iters); + for (int i = 0; i < iters; ++i) { + const auto t0 = clock::now(); + for (int b = 0; b < batch; ++b) { + fn(); + } + const auto t1 = clock::now(); + const double batch_ns = std::chrono::duration(t1 - t0).count(); + samples_ns.push_back(batch_ns / batch); + } + + std::sort(samples_ns.begin(), samples_ns.end()); + TimingResult result; + result.min_ns = samples_ns.front(); + result.median_ns = samples_ns[samples_ns.size() / 2]; + return result; +} + +inline double bytes_per_sec(size_t bytes, double ns) { + if (ns <= 0.0) { + return 0.0; + } + return static_cast(bytes) / (ns * 1e-9); +} + +inline double gflops(double flops, double ns) { + if (ns <= 0.0) { + return 0.0; + } + return flops / (ns * 1e-9) / 1e9; +} diff --git a/apps/ggml/vcpkg-configuration.json b/apps/ggml/vcpkg-configuration.json new file mode 100644 index 000000000000..a0daf83101ff --- /dev/null +++ b/apps/ggml/vcpkg-configuration.json @@ -0,0 +1,5 @@ +{ + "overlay-ports": [ + "../vcpkg/ports" + ] +} diff --git a/apps/ggml/vcpkg.json b/apps/ggml/vcpkg.json new file mode 100644 index 000000000000..bac5a1fefe0f --- /dev/null +++ b/apps/ggml/vcpkg.json @@ -0,0 +1,9 @@ +{ + "name": "halide-ggml-app", + "version": "22.0.0", + "license": "MIT", + "builtin-baseline": "66c0373dc7fca549e5803087b9487edfe3aca0a1", + "dependencies": [ + "ggml" + ] +} diff --git a/apps/vcpkg.json b/apps/vcpkg.json index 1c3c5805f7ca..381ba6e211ac 100644 --- a/apps/vcpkg.json +++ b/apps/vcpkg.json @@ -10,6 +10,7 @@ "platform": "(windows & x64 & !uwp & !xbox) | (linux & x64) | (linux & arm64)" }, "eigen3", + "ggml", "libjpeg-turbo", "libpng", "onnx", diff --git a/apps/vcpkg/ports/ggml/portfile.cmake b/apps/vcpkg/ports/ggml/portfile.cmake new file mode 100644 index 000000000000..942b99f5f5b4 --- /dev/null +++ b/apps/vcpkg/ports/ggml/portfile.cmake @@ -0,0 +1,26 @@ +vcpkg_check_linkage(ONLY_STATIC_LIBRARY) + +vcpkg_from_github( + OUT_SOURCE_PATH SOURCE_PATH + REPO ggml-org/ggml + REF eced84c86f8b012c752c016f7fe789adea168e1e # v0.15.3 + SHA512 3295c064aff295b0387249d5dec7860b620de82c8361197888df186be18270ede253ab7bce3358b1fb3020f11d01f0f8a29f7d268bff44666f8a8f3ea832781e +) + +# We set GGML_BACKEND_DL=OFF to keep the CPU backend linked, not dlopen'ed, +# because apps/ggml needs ggml-cpu's internal symbols at link time. +vcpkg_cmake_configure( + SOURCE_PATH "${SOURCE_PATH}" + OPTIONS + -DBUILD_SHARED_LIBS=OFF + -DGGML_BACKEND_DL=OFF + -DGGML_BUILD_TESTS=OFF + -DGGML_BUILD_EXAMPLES=OFF +) + +vcpkg_cmake_install() +vcpkg_cmake_config_fixup(PACKAGE_NAME ggml CONFIG_PATH lib/cmake/ggml) + +vcpkg_install_copyright(FILE_LIST "${SOURCE_PATH}/LICENSE" "${SOURCE_PATH}/AUTHORS") + +file(REMOVE_RECURSE "${CURRENT_PACKAGES_DIR}/debug/include" "${CURRENT_PACKAGES_DIR}/debug/share") diff --git a/apps/vcpkg/ports/ggml/vcpkg.json b/apps/vcpkg/ports/ggml/vcpkg.json new file mode 100644 index 000000000000..043c2b93892d --- /dev/null +++ b/apps/vcpkg/ports/ggml/vcpkg.json @@ -0,0 +1,17 @@ +{ + "name": "ggml", + "version": "0.15.3", + "description": "Tensor library for machine learning", + "homepage": "https://github.com/ggml-org/ggml", + "license": "MIT", + "dependencies": [ + { + "name": "vcpkg-cmake", + "host": true + }, + { + "name": "vcpkg-cmake-config", + "host": true + } + ] +} diff --git a/doc/Approximation.md b/doc/Approximation.md new file mode 100644 index 000000000000..2c2964c5f5ff --- /dev/null +++ b/doc/Approximation.md @@ -0,0 +1,936 @@ +# Approximation: a core concept for lossy, quantified Func substitution + +This is a design document: it explains the concepts and the reasoning behind +them. "Summary of decisions and open items" at the end lists what is still +undecided. + +## Motivation + +Two pieces of prior work motivate this: + +1. An out-of-tree GGML application hand-implements ~24 quantized weight formats + as pairs of Halide Generators (quantize, dequantize). Each file independently + encodes its own byte layout, scale/bias math, and (for K-quants and a few + others) a call to GGML's own reference quantizer. Every type duplicates the + same shape of logic (block layout, scale search, bit-packing) by hand, and + the "quantize happens once offline, dequantize happens on every inference + call" relationship between the two directions is enforced by nothing but + convention and comments. + +2. A private Python research prototype builds the same formats compositionally: + a small `Approximation` ABC (`encode`/`decode`, each operating on Halide + `Func`s) with a handful of primitives (block reshaping, a linear integer + quantizer, a shift-by-min helper, bit-packers) that compose to reconstruct + the K-quant family. It is Python-only and JIT-only, and covers only + quantize/dequantize round trips. + +This document describes the C++ realization of that idea as a first-class Halide +concept, `Approximation`, plus the surrounding API needed to wire one into a +real pipeline: `Func::approximate_by()`, which splices an approximation's round +trip into an existing call graph; `Pipeline::sever()`, which optionally splits +the result across a compile-time boundary; and +`tools/halide_approximation_testing.h`, which checks what an approximation +claims. + +## Why this can't just be ordinary scheduling + +Halide's algorithm/schedule separation depends on schedule directives being +meaning-preserving: `.compute_root()` vs `.compute_at()` never changes what a +pipeline computes, only how. An `Approximation` is the opposite by design: it +deliberately changes the *value* computed (a real weight becomes a +quantized-then-dequantized approximation of itself), in a bounded, quantified +way. Wiring that in by disguising it as an ordinary Func substitution (e.g. a +custom `.in()` wrapper with no other marking) would make a semantics-changing +operation look, to any future reader, like a semantics-preserving one. It needs +its own footing in the API. + +## Core concept: `Approximation` + +An `Approximation` is a value-semantic, type-erased handle (in the style of +`std::function`) to a lossy transformation of one or more Funcs' values: +`decode(encode(f))` approximately reproduces `f`. It operates purely on `Func`s +and makes no claim about *where* or *when* either half is computed (see "Scope: +placement is not semantics"). + +```cpp +class Approximation { +public: + Approximation(); // undefined + template Approximation(T &&unit); // duck-typed, implicit + template Approximation(T &&unit, std::string label); + + bool defined() const; + bool same_as(const Approximation &other) const; + Approximation labelled(std::string label) const; + std::string label() const; + + EncodeResult encode(const std::vector &inputs, + const ApproximationPorts &input_ports = {}) const; + DecodeResult decode(const std::vector &encoded, + const ApproximationPorts &input_ports = {}) const; + + ApproximationSignature signature(const ApproximationPorts &inputs = {}) const; + Func error_bound(const std::vector &inputs, + const std::vector &encoded) const; + bool lossless() const; + std::string describe(const ApproximationPorts &inputs = {}) const; +}; +``` + +The handle's `encode()`/`decode()` always take and return *vectors* of Funcs, +even though a plain quantizer only uses one. That is what makes the combinators +possible: an inner stage can produce several Funcs (a codes Func and a separate +scale Func), and the next stage needs to consume all of them, or pick one to act +on. There is no base class; the type-erasure is what lets `Compose` and +`Parallel` hold a runtime-heterogeneous list of stages. + +**Identity.** A copy of a handle is the *same* stage (`same_as()` is true); +converting a plain unit to an `Approximation` twice makes two *distinct* stages, +even if the units compare equal. Every stage invoked through a handle, including +from inside another unit's `encode`/`decode`, is recorded in the result, so a +caller finds a stage's Funcs by keeping a handle to it and passing copies of +that handle to the combinators (see "Introspection"). + +**Labels.** Each handle has a label for display: `labelled("x")` (or the +two-argument constructor) sets it, else the unit's `std::string name() const` if +it has one, else the unit's type name with `Halide::` qualifiers stripped +(`Compose`, `LittleEndianScalarPack`). The label lives in state +shared by all copies of the handle, so `labelled()` affects every copy and +returns a handle that is `same_as` the original, even though the method is +`const`. + +### Defining a unit + +A unit is any type with const-callable `encode` and `decode` methods. Each +direction independently takes one of three forms: + +| Form | Signature | +| ---------- | --------------------------------------------------------------------------------------- | +| single | `Func encode(const Func &) const` | +| multi | `std::vector encode(const std::vector &) const` | +| port-aware | `std::vector encode(const std::vector &, const ApproximationPorts &) const` | + +`decode` is the same with `encoded` in place of `inputs`. The forms may be mixed +across directions. A single-form direction handed any other number of Funcs is +an error. If a type offers both a vector and a `Func` overload for one +direction, the vector form is used; the port-aware form is preferred over both. +Types whose return types do not match, or with only one direction, do not +convert. The handle stores a decayed copy of the unit and only calls const +methods, so units needing mutable state must hold it in `mutable` members or +behind a pointer. + +A unit only *returns its outputs*. It does not declare or register its +intermediate Funcs: after the unit returns, the framework walks the outputs' +definitions (pure, update, and extern-argument references), stopping at the +stage's inputs, and reports every other Func it reaches in topological order +(producers before consumers, ties broken by name). That includes pure-only +Funcs, and Funcs the unit references from outside (e.g. a shared lookup table); +callers filter. A unit that calls other `Approximation` handles inside its own +`encode`/`decode` (as `Compose` does) needs no extra bookkeeping: those calls +are traced automatically. + +The port-aware form additionally receives the ports (names and constraints) of +the encode-side inputs: `encode(inputs, input_ports)` and +`decode(encoded, input_ports)`. It exists so that combinators can route by port +name and forward context to their children; leaf units rarely need it. + +A unit may additionally provide any of the following optional members. + +- `ApproximationSignature signature() const` (static: the same ports whatever + the context) *or* `signature(const ApproximationPorts &inputs) const` + (contextual: given the resolved input ports, empty if unknown, return the full + signature; used by combinators and shape-polymorphic units). +- `std::string name() const`: the default label. +- `std::vector children() const`, and optionally + `std::vector child_inputs(const ApproximationPorts &inputs) const` + (the ports each child would receive, parallel to `children()`): the stages + this unit is built from, used by `describe()` and `check_ranges()`. Only the + unit knows how it routes ports to children, so this cannot be derived. +- `Func error_bound(const std::vector &inputs, const std::vector &encoded) const`: + a per-element bound on `abs(decode(encode(x)) - x)`, as a Func over + `inputs[0]`'s arguments. +- `bool lossless() const`: the round trip is exact (a bound of zero). + +If a unit provides both signature forms, the contextual one is used. A unit with +only an `error_bound()` is not `lossless()` even when the bound happens to be +zero. Both accuracy declarations are *claims that hold when every input port's +declared range (its precondition) is respected*. They are never enforced at run +time; see "Verification" for how they are checked. + +`Approximation::error_bound(inputs, encoded)` returns the unit's bound, else a +zero-valued Func if the unit is `lossless()`, else an undefined Func. + +### Worked example + +Elementwise units need no struct: `Pointwise{name, encode_fn, decode_fn}` builds +one from a pair of `Expr -> Expr` lambdas (or `std::vector` ones for +Tuple-valued Funcs; pass lambdas with concrete parameter types, not `auto`). A +cast or an offset is one inline `Pointwise`; the optional `with_types`, +`with_ranges` and `with_lossless` declare (and let tests check) more than its +arity. + +```cpp +// Drops the low bit of an integer: lossy, with error at most 1. +Approximation drop_lsb = + Pointwise{"drop_lsb", + [](Expr x) { return x >> 1; }, + [](Expr x) { return x << 1; }} + .with_error_bound([](Expr) { return Expr(1); }); + +// Signed 4-bit codes to stored nibbles: exact for codes in [-8, 7]. +Approximation offset = + Pointwise{"offset", + [](Expr x) { return cast(x + 8); }, + [](Expr x) { return cast(cast(x) - 8); }} + .with_types(Int(8), UInt(8)) + .with_ranges(ApproximationRange(-8, 7), ApproximationRange(0, 15)) + .with_lossless(); +``` + +Anything less regular is a small struct. This int8 quantizer with one scale per +block uses the multi form, declares a signature with a range on its codes, and +declares an error bound. Its `amax` reduction is found by the framework; the +unit never mentions it. + +```cpp +struct BlockQ8 { + static constexpr int kBlock = 32; + + std::vector encode(const std::vector &in) const { + Func x = in[0]; + Var i("i"), b("b"); + RDom r(0, kBlock, "r"); + Func amax("amax"); + amax(b) = 0.0f; + amax(b) = max(amax(b), abs(x(b * kBlock + r))); + Func scale("scale"); + scale(b) = amax(b) / 127.0f; + Func codes("codes"); + Expr inv = select(scale(i / kBlock) != 0.0f, 1.0f / scale(i / kBlock), 0.0f); + codes(i) = cast(clamp(round(x(i) * inv), -127, 127)); + return {codes, scale}; + } + + std::vector decode(const std::vector &enc) const { + Var i("i"); + Func out("dequantized"); + out(i) = cast(enc[0](i)) * enc[1](i / kBlock); + return {out}; + } + + // Written in the encode direction: decode consumes `outputs` and + // produces `inputs`. + ApproximationSignature signature() const { + return {{{"values", Float(32), 1}}, + {{"codes", Int(8), 1, ApproximationRange(-127, 127)}, + {"scale", Float(32), 1}}}; + } + + // Half a step, plus a little slack for float rounding. + Func error_bound(const std::vector &, const std::vector &encoded) const { + Var i("i"); + Func bound("bound"); + bound(i) = abs(encoded[1](i / kBlock)) * 0.5001f; + return bound; + } +}; + +Approximation q = BlockQ8{}; +``` + +## Composition + +### Core units + +`Approximation.h` provides the combinators and a few generic leaf units for +layout and bit packing. All are plain structs that convert to `Approximation`, +and all declare signatures. Domain-specific quantizers (block scales, rounding +policies, code alphabets) are not part of the core: they live in client code, +like `BlockQ8` above or the Q4_0 quantizer in lesson 25 (GGML's own quantizers +are defined in the GGML app). + +| Unit | Description | +| ---------------------------------------------------------- | --------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `Compose{stages...}` | Sequential composition, in encode order: `encode` runs the stages first to last, `decode` runs them last to first. The last stage's encoded output is the Compose's own. Lossless iff all stages are; no error bound is declared for lossy compositions, because bounds do not compose without knowing how errors propagate. | +| `Parallel{children...}` / `Parallel{{"port", child}, ...}` | Product: applies each child to its own share of the Funcs. See below. | +| `TrustedInverse{encoder, decoder}` | Takes `encode` from one approximation and `decode` from another. The escape hatch out of `Compose`'s structural guarantee (see below). | +| `Choose{cond, if_true, if_false}` | Keeps whichever handle `cond` selects at construction, so it can be found by that handle. | +| `Identity{}`, `Permute{permutation}` | Pass Funcs (and their port names) through unchanged / reordered. Both are lossless. | +| `Pointwise{...}` | Elementwise `out(vs) = fn(in(vs))`; the output Funcs are named `name + "_encode"` and `name + "_decode"` (the four-argument constructors choose both names and the pure Var prefix). `with_types`, `with_ranges`, `with_lossless` and `with_error_bound` declare the signature, precondition and guarantee ranges, losslessness and an error bound. | +| `BlockReshape` | Flat row to fixed-size records; lossless. | +| `StructLayout` | Logical Funcs to a struct-typed record Func, one field each; lossless. | +| `LittleEndianScalarPack` | A word per record to and from a leading byte dimension; lossless. | +| `PlanarFieldPack` | Fixed-width fields packed into bytes; lossless for inputs in its declared range. | + +A four-bit scheme, built from these and a quantizer of your own (`quant`, here +one whose codes are declared to lie in `[-8, 7]`), reads in the order `encode` +runs it: reshape to blocks, quantize, then shift the codes to `[0, 15]` (the +`offset` above) and pack them. + +```cpp +Approximation pack = PlanarFieldPack{4, 8}; +Approximation scheme = Compose{ + BlockReshape{16}, quant, + Parallel{{"codes", Compose{offset, pack}}}}; +``` + +**`Compose` and `TrustedInverse`.** Every `Approximation` is meant to be an +approximate identity factored into a `decode`-after-`encode` pair. `Compose` +preserves that structurally: it interleaves its stages' `encode`s and `decode`s +in mirror order, so both halves provably come from one stage list. +`TrustedInverse` pairs an `encode` and a `decode` from unrelated approximations, +so nothing structural guarantees they compose to an identity; the caller is +*trusted* to have supplied a true inverse pair. The motivating case is a scheme +whose forward map is an opaque offline black box (a per-block codeword search, +typically an extern call) that no composition of Funcs reproduces bit-for-bit, +but whose reverse map is an ordinary `Compose`. The unused half of each side is +never called. + +**`Parallel`.** A product combinator with two forms; they cannot be mixed. + +- *Positional*: `Parallel{a, b, ...}`. Child `i` gets a consecutive slice of the + Funcs. On encode its width is the number of inputs of its signature, on decode + the number of outputs (its encoded ports); one if the signature is unknown. + The slices must exactly cover the Funcs, or it is an error stating both + counts. `Identity{}` passes one Func through. +- *Named*: `Parallel{{"codes", a}, {"scale", b}}`. Each entry routes the one + port of that name to its child. Ports not mentioned pass through unchanged, in + place. A child's outputs replace its port in place, so a child may expand a + port into several on encode; the mirror collapses them on decode, where the + child must yield exactly one Func. A name that is missing or ambiguous, or + routed twice, is an error. Without input ports (no context) the signature is + unknown. + +`Parallel` is lossless iff all children are and declares no error bound. Its +children appear in `describe()` and `check_ranges()`, and are traced like any +other combinator's. + +### Ports and naming + +Every Func a stage consumes or produces is a *port*: + +```cpp +struct ApproximationPort { + std::string name; + std::optional type; + std::optional dimensions; + std::optional range; // constant [lo, hi], see below +}; +using ApproximationPorts = std::vector; +``` + +`type` and `dimensions`, when set, are checked against the actual Func at run +time; a port with a Tuple-valued Func never has a `type`. `range` is a declared +bound on the port's *values*, never checked on the normal encode/decode path +(see "Verification"). + +A port name identifies one wire, in both directions. Declared names are +*checked, not substituted*. + +- **Names flow in.** The names of the input ports come from, in order: the ports + the caller (or the upstream stage) passed, else the unit's declared signature + inputs, else positional `"0"`, `"1"`, .... A declared input name is only a + default, used when no name flows in; it never renames a flowing wire. Types + and dimensions are still checked against the declared ones. +- **Output ports** are the declared signature's outputs (their count must match + the number of Funcs the unit returned) or, for an undeclared unit or one whose + signature is unknown, follow the *naming rule*: if the output count equals the + input count, output `i` takes input `i`'s name (so a single-Func unit + preserves its input's name); otherwise the outputs are named positionally. +- The resolved output ports, with unset types and dimensions filled in from the + actual Funcs, are returned as `EncodeResult::encoded_ports` / + `DecodeResult::decoded_ports`, ready to hand to the next stage. +- **Mirror invariant.** For every stage, decode output `i` has the name of + encode input `i`, and decode inputs have the names of the encode outputs. So a + by-name `Parallel` routes the same port names in both directions, even across + `StructLayout`, and decode ports never collide. +- **Stand-alone decode.** `decode(encoded, input_ports)` takes the encode-side + context, computed statically: with none given, it is threaded from the root's + default input names via signatures, exactly as `describe()` does. So a decode + run on its own (e.g. after `sever` severs the encode) names its ports as an + encode+decode would. +- `Func::approximate_by()` passes no names: the root's declared input names (or + `"0"`) are the defaults. +- **Limit.** If a stage in a `Compose` has an unknown signature (an undeclared + multi-Func unit), the contexts after it are unknown and fall back to defaults, + so mirror naming past it is best-effort. + +### Signatures and validation + +```cpp +struct ApproximationSignature { + ApproximationPorts inputs, outputs; // encode direction + bool known = true; + static ApproximationSignature unknown(ApproximationPorts inputs = {}); +}; +``` + +The signature is written in the encode direction: `encode` consumes `inputs` and +produces `outputs`; `decode` consumes `outputs` and produces `inputs`. An +*undeclared* unit has no signature of its own, but +`Approximation::signature(inputs)` still resolves what it can without running +anything: + +- a static signature is returned as is; +- a contextual signature is computed from `inputs`; +- an undeclared unit with a single-Func `encode` has one input (`inputs`, or + `"0"` if none were given) and one output with the same name; +- anything else (an undeclared multi-Func unit) is *unknown* (`known == false`; + `inputs` echoes the context and `outputs` is empty). + +Combinators derive their signatures from their children's, and are unknown +whenever a child is: `Compose` chains its stages' signatures in encode order, +`Parallel` splices its children's ports into the slices they handle, +`TrustedInverse` and `Choose` report the encoder's and the chosen stage's, +`Identity` echoes its inputs (unknown if there are none), and `Permute` permutes +them. + +On every `encode`/`decode` call, the handle checks each input Func against its +port's `type` and `dimensions` (both the ports the caller gave and the unit's +declared ones), and each output Func against the declared output ports. A +mismatch is a `user_error` naming the stage, the direction, the port, and +expected vs actual. Declared signatures are how a unit's Func-level contract +(e.g. "one float32 Func of dimension 1") gets checked at the point where a +scheme is assembled, rather than deep inside Halide's own lowering. + +## Introspection + +**Intermediates.** Each call to `encode`/`decode` reports, in `intermediates`, +every Func reachable from the stage's outputs without passing through one of its +inputs, excluding the inputs and outputs themselves (see "Defining a unit"). The +Funcs with update definitions among them need scheduling by whoever calls +`encode`; see `approximate_by` below. + +**Trace.** The outermost handle call on a thread opens a trace, and every handle +call nested inside it appends a record, children first, then the stage itself. +Encode and decode share one trace (a `decode` called from inside an `encode` +appears in that `encode`'s trace). The trace is a tree, + +```cpp +struct ApproximationTraceNode { + Approximation stage; + std::string label; // stage.label() when the call finished + std::vector ports; // the stage's outputs + std::vector port_names; // parallel to `ports` + std::vector intermediates; // discovered for this stage alone + std::vector children; // in invocation order + std::vector inputs; // what the call consumed + std::vector input_names; // parallel to `inputs` +}; +``` + +exposed as `EncodeResult::trace`, `DecodeResult::trace`, and +`ApproximationResult::encode_trace`/`decode_trace`. The flat `stage_outputs` +lists (`ApproximationStageOutputs{stage, ports, port_names, intermediates}`) are +its post-order flattening. `operator<<` prints a trace node, or a whole +`ApproximationResult` (under `encode:` and `decode:`), as an indented tree of +`label -> port=Func`, with each stage's intermediates on an `intermediates:` +line: + +``` +encode: + BlockQ8 -> codes=codes, scale=scale + intermediates: amax +decode: + BlockQ8 -> values=dequantized +``` + +**Looking up a stage.** Stages are found by handle, not by position or path: + +```cpp +Func encoded_by(const Approximation &stage, size_t port = 0) const; +Func decoded_by(const Approximation &stage, size_t port = 0) const; +Func encoded_by(const Approximation &stage, const std::string &port) const; +Func decoded_by(const Approximation &stage, const std::string &port) const; +``` + +on `ApproximationResult`. It is an error if `stage` was not invoked in that +direction, or was invoked more than once (the lookup would be ambiguous); an +out-of-range positional port yields an undefined Func, and an unknown or +duplicated port name is an error whose message lists the ports the stage has. + +```cpp +Approximation qh = LittleEndianScalarPack{}; +Compose scheme{BlockReshape{32}, qh}; +ApproximationResult r = f.approximate_by(scheme, {g}); +Func bytes = r.decoded_by(qh); +``` + +**Stage ports.** `ApproximationResult::stage_ports()` returns every Func that is +an output port of some stage in either direction (deduplicated by name, encode +side first, each side in post-order), excluding `replacement`; +`is_stage_port(f)` tests membership. Callers use it to schedule stage boundaries +alongside reductions, e.g. +`if (f.has_update_definition() || r.is_stage_port(f)) f.compute_root();`. + +**`describe()`.** `Approximation::describe(inputs = {})` (also `operator<<`) +renders a stage's structure without running anything: one +`label (inputs) -> (outputs)` line per stage, where a port prints as +`name: type xN in [lo, hi]` (unset parts are left out), followed by the stage's +children (`Compose`: encode order; `Parallel`: its children; `TrustedInverse`: +encoder then decoder; `Choose`: the chosen stage), indented by two spaces and +given the ports they would receive. For the four-bit scheme above (with the +quantizer from lesson 25): + +``` +Compose (values x1) -> (bytes: uint8 x2, scale: float32 x1) + BlockReshape (values x1) -> (blocks x2) + Q4_0Quantizer (blocks: float32 x2) -> (codes: int8 x2 in [-8, 7], scale: float32 x1) + Parallel (codes: int8 x2 in [-8, 7], scale: float32 x1) -> (bytes: uint8 x2, scale: float32 x1) + Compose (codes: int8 x2 in [-8, 7]) -> (bytes: uint8 x2) + offset (codes: int8 x2 in [-8, 7]) -> (codes: uint8 x2 in [0, 15]) + PlanarFieldPack (codes x2 in [0, 15]) -> (bytes: uint8 x2) +``` + +An unknown signature prints as `(unknown signature)`. Where a stage's context is +known, each input whose declared range is not guaranteed by its producer is +flagged on the following line, e.g. +`! input 'fields': [0, 16] not within [0, 15]` (or `not guaranteed` when the +producer declares no range). + +## `approximate_by`: wiring an `Approximation` into a call graph + +```cpp +ApproximationResult Func::approximate_by(const Approximation &p, + const std::vector &consumers); + +struct ApproximationResult { + Func replacement; // decode's round-trip output; already + // spliced into every Func in `consumers` + std::vector encoded; // the Funcs encode() produced + ApproximationPorts encoded_ports; // parallel to `encoded` + std::vector intermediates; + std::vector encoded_stage_outputs, + decoded_stage_outputs; + ApproximationTraceNode encode_trace, decode_trace; + // encoded_by(), decoded_by(), stage_ports(), is_stage_port(): see above +}; +``` + +`approximate_by` runs `p.encode({*this})` and `p.decode(encoded)`, requires the +decode to yield exactly one Func whose dimensionality and types match `*this` +(*the signature contract*: `decode(encode(f))` reproduces `f`'s arg list and +value type exactly, which is what makes it valid to splice back in), and then +replaces every call to `*this` inside each Func of `consumers` with a call to +that Func. A Func cannot be its own consumer. `intermediates` is `encoded`, then +the encode side's discovered intermediates, then the decode side's, without +duplicates, and never contains the original Func or `replacement`. Like all +discovered intermediates it may include Funcs the units referenced from outside. + +The contract is not enforced generically at the `Approximation` level: each +concrete unit is responsible for it, and round-trip error and convergence are +treated as testable properties (see "Verification") rather than type-level +guarantees. `approximate_by`'s check at the point of substitution catches +violations of the shape contract, just not at definition time. + +### Why not `Func::in` + +The targeted form `g.in(f)` is eager: it rewrites `f` to call a new wrapper +immediately, so the graph reflects the change as soon as the call returns. But +`in()` always substitutes an *identity* wrapper; `approximate_by` must +substitute a different computation, `decode(encode(f))`. The global form +`f.in()` is not an option either, because it is deferred: it registers a wrapper +that `wrap_func_calls` applies during `lower()`, after +`configure()`/`generate()`/`schedule()` have run. Anything that reasons about +what a consumer actually calls before then (in particular, `sever`'s +`configure()`-time split) would see stale, pre-substitution state. + +### The mechanism: eager and destructive, like `rfactor` + +`Stage::rfactor` is the right precedent. It never defers to a lowering pass: it +builds a new `Func`, calls `define_update` on it immediately, and rewrites the +original Function's own definition, all synchronously, as part of the +`.rfactor()` call itself. By the time it returns, the graph already reflects the +change, which is why a caller can immediately turn around and schedule the new +Func. + +`approximate_by` behaves the same way, using the same class of primitive +`WrapCalls.cpp` already relies on internally, +`Function::substitute_calls(orig, substitute)`, but invoked immediately, on an +explicitly given set of consumers, instead of registered for a later pass. It is +a member of `Func` (`f.approximate_by(p, consumers)`), not a free function, +because it is a graph-editing operation on `f` in exactly the sense that +`f.in(...)`, `f.clone_in(...)` and `f.rfactor(...)` are. No new internal +primitive was needed; `Function::substitute_calls` is already an ordinary +method, and `approximate_by` lives inside libHalide. + +**Reporting `intermediates` is not optional.** Both `encode` and `decode` can +introduce Funcs with update definitions (per-block reductions, a shift-by-min +helper's own min-reduction). Left unscheduled, Halide computes them at the +innermost valid loop level by default. But that is only a default: the caller +has no way to override it, or to apply the fusion patterns from "Scope: +placement is not semantics" (e.g. `compute_at`-ing `encoded` into a producer for +dynamic activation requantization), unless it has the Funcs in hand. Because +units cannot be trusted to declare them, the framework discovers them (see +"Introspection") and bundles them into `intermediates` so the caller can +schedule all of them, not just the primary output. + +### Consequence: consumers must already exist + +Because the substitution is eager, `approximate_by` can only rewrite Funcs that +are already built at the point of the call. There is no equivalent of the global +`.in()` (redirect *every* current and future consumer). This is a real +capability loss, but it is the same scoping `rfactor` lives with, and it matches +how Generator code is written: within `generate()` (or `configure()`), `f` and +its consumers are typically built together, so passing `consumers` explicitly +costs nothing. It only forecloses transparently intercepting calls inside a +large, opaque, externally-authored algorithm whose call sites cannot be +enumerated, which is out of scope. + +## Scope: placement is not semantics + +An `Approximation`'s `encode`/`decode` never make any claim about *where* or +*when* they are computed relative to the rest of the pipeline. An early draft +proposed otherwise, that `encode`'s output could always be treated as "the +offline half", and was rejected on a concrete counterexample: **dynamic +activation requantization**. + +- **Static weight quantization**: `encode` (e.g. Q4_0 quantize) runs exactly + once, ever, fully decoupled from any inference call, a genuine + compile-time/compilation-unit boundary. `decode` is fused inline into the + consumer's inner loop and never materialized as its own Func. +- **Dynamic activation requantization**: `encode` (quantize a just-computed + activation tile) needs to be fused into the *producer's* schedule: same + granularity, same loop nest, no separate storage, recomputed every call. + `decode` is fused into the consumer's tiles exactly as before. + +Same `Approximation`, opposite treatment of where `encode` is computed. If +"encode implies offline" were baked into the interface, the activation case +would need an escape hatch to override it, at which point the shortcut has +bought nothing; and it would invite tooling to assume every quantize step is +safe to hoist to conversion time, a correctness trap for anything computed at +inference time. + +**Consequence, and a scope reduction**: fusing `encode` into a producer or +`decode` into a consumer needs no new Halide feature. Ordinary `.compute_at()` / +`.compute_inline()` on `ApproximationResult`'s `replacement`, `encoded`, +`intermediates` and `stage_ports()` already achieves it, since they are regular +Funcs in the call graph. `sever` (below) is needed only for the strictly +narrower case of actually severing the graph into two separately-compiled +artifacts, the static-weight case. + +## `sever`: scope + +`Pipeline::sever()` is deliberately independent of `Approximation`: it operates +on Funcs. + +```cpp +SeverResult sever(const std::vector &to_sever); +SeverResult sever(const std::vector &to_sever, + const std::vector &bind_to); +SeverResult sever(const std::vector &to_sever, + const std::vector &names); + +struct SeverResult { + Pipeline offline; // computes to_sever's true values + std::vector online_inputs; // one per to_sever, same order +}; +``` + +It rewrites every call to each Func in `to_sever`, anywhere in the pipeline's +transitive call graph, to call an `ImageParam` of matching type and +dimensionality instead (a fresh one, one named by `names`, or the caller's +`bind_to`). Anything reachable only from `to_sever` (a per-block reduction, say) +belongs to the offline half and keeps its true computation. Like `rfactor` it is +eager and destructive. `offline` is a `Pipeline` whose outputs are `to_sever`; +realize it once (JIT), or compile it as its own artifact (AOT), and feed the +result to `online_inputs` before realizing the original pipeline. + +**Decision: "seam exposure," not automatic pipeline splitting.** Given Funcs +that should become a compile-time boundary, the result is: + +- their computation exists in one compile, as ordinary outputs (the "offline" + artifact); +- same-shaped inputs exist in another compile, standing in for them (the + "online" pipeline); +- both are ordinary, statically declared Generator I/O; there is no dynamic + discovery of new ports mid-`generate()`. + +True automatic splitting (one Generator definition, two artifacts emitted +automatically, no extra static I/O declared by the author) was considered and +rejected: it would require a Generator to discover an extra Input/Output +*during* `generate()`, based on the structure of a Func graph that does not yet +exist when `configure()` declares I/O. That is a phase-ordering problem, not +just an ergonomics one. + +**Restriction:** each severed Func must be single-valued (no Tuples). + +### Composability with `approximate_by` + +Both operations are eager and destructive, so they compose in program order. The +graph state at every point *is* the true state; no later lowering pass can +silently change what a Func calls out from under code that already ran. The +`encoded` Funcs of an `ApproximationResult` are exactly what to hand to `sever`: +they are the Funcs the consumers' rewritten call graph depends on. + +```cpp +ApproximationResult r = f.approximate_by(scheme, {consumer}); +SeverResult split = Pipeline({consumer}).sever(r.encoded); +``` + +A Generator authoring an op from scratch usually does not need `approximate_by` +at all: it can build `encode`/`decode` itself and use `decoded[0]` wherever the +math needs the value, since it is writing the consumer fresh anyway. +`approximate_by` earns its keep when the consumer already exists as written code +the author does not want to edit by hand. + +### Non-goal: provenance checking + +Nothing here guarantees that the `Approximation` used to produce the offline +artifact in one compile is *actually* the same one the online compile expects +when decoding it. Correctness rests on both sides building the same scheme. +Embedding a scheme fingerprint in the artifact and checking it at load time is +deferred; this is a user obligation, as any hand-written quantize/dequantize +split already does. + +## Generator shape + +Generators already support dynamic I/O declared before `generate()` runs: +`configure()` exists so that `add_input<>()`/`add_output<>()` can be called +based on `GeneratorParam` values decided earlier. Two additions let +`configure()` adopt the halves of a `sever` split as ports: + +```cpp +template> GeneratorInput *add_input(const ImageParam &existing); +template> GeneratorOutput *add_output(const Func &existing); +``` + +Each declares a port backed directly by the existing object and named after it; +`T` may be given explicitly (`add_input>(q_in)`) to check the +object's type and dimensionality and give the stub statically typed buffers. +Both may only be called from `configure()`. + +This lets one `configure()` build a whole round trip, split it, and adopt +whichever half applies, leaving `generate()` an empty stub. The quantize and +dequantize Generators share the body; only the ports differ: + +```cpp +enum class Direction { Quantize, Dequantize }; + +class Codec : public Generator { +public: + GeneratorParam direction{ + "direction", Direction::Quantize, + {{"quantize", Direction::Quantize}, {"dequantize", Direction::Dequantize}}}; + + void configure() { + Approximation scheme = BlockQ8{}; // e.g. selected from GeneratorParams + + ImageParam x(Float(32), 1, "x"); + Var i("i"); + Func y("y"); + y(i) = x(i); // stands in for whatever consumes the value + + ApproximationResult r = Func(x).approximate_by(scheme, {y}); + for (Func f : r.intermediates) { + if (f.has_update_definition() || r.is_stage_port(f)) { + f.compute_root(); + } + } + + std::vector names; + for (const ApproximationPort &p : r.encoded_ports) { + names.push_back(p.name + "_in"); + } + SeverResult split = + Pipeline({y}).sever(r.encoded, names); + + if (direction == Direction::Quantize) { + add_input(x); + for (Func e : split.offline.outputs()) { + add_output(e); + } + } else { + for (const ImageParam &in : split.online_inputs) { + add_input(in); + } + add_output(y); + } + } + + void generate() {} +}; +``` + +The encoded form's arity and layout are whatever the scheme chose, so the +Generator's public signature is scheme-dependent. That follows from a layout +choice each `Approximation` makes, and the framework does not paper over it: + +- *Packed*: one opaque byte buffer (or one struct-typed Func), fields recovered + inside `decode`. +- *Planar*: multiple typed Funcs (a `float16` delta Func, an `int8` quants + Func), more Halide-native and type-safe, but the public signature grows with + the scheme's field count. + +Known rough edge, accepted for now: `add_input`/`add_output` return raw +pointers, so a Generator that wants to keep them needs member-pointer +bookkeeping that the static `Input<>`/`Output<>` member style does not. + +## Verification + +`tools/halide_approximation_testing.h` (namespace +`Halide::ApproximationTesting`, header-only, like `halide_image_io.h`; not part +of `Halide.h`) checks what a scheme claims. It is a test-time helper: nothing is +added to generated pipelines. Everything that runs Halide code uses the JIT, +with every stage boundary `compute_root`'d and traced so its values can be read +back; it does not change how a pipeline you build yourself is scheduled. + +**Declared facts.** + +- `ApproximationPort::range = ApproximationRange{lo, hi}`: constant `double` + bounds (integers are exact up to 2^53). On an encode *input* port it is a + *precondition*: the range the stage requires for its declared properties to + hold (`PlanarFieldPack` with 4-bit fields needs `[0, 15]`). On an encode + *output* port it is a *guarantee* (symmetric int8 codes lie in `[-127, 127]`). +- `Approximation::describe()` and `check_ranges(a, inputs = {})` compare, for + every stage whose input context is known, each producer's guaranteed range + with the consumer's precondition, without running anything. `check_ranges` + returns a diagnostic per failure, e.g. + `path: input 'p': [0, 15] not within [0, 7]`, or + `path: input 'p': requires [0, 7], but the producer declares no range`; an + empty result means every declared precondition is statically guaranteed. + Ranges only flow through units that declare them, so unknown is common; + unknown is reported, never assumed satisfied. +- `error_bound()` and `lossless()` (see "Defining a unit") declare accuracy. + `BlockQ8` above declares half a step; a unit that declares none is treated as + having no bound. + +Ranges are never enforced by `encode`/`decode`, so they cost nothing in +generated code. + +**Distributions.** A `Distribution` is a plain value describing how to fill a +buffer; `generate(dist, type, extents, seed)` draws it with a built-in +xoshiro256\*\* PRNG using only integer arithmetic and IEEE `+ - *`, so data is +identical on every platform (`normal()` is therefore an Irwin-Hall +approximation, not Box-Muller). Elements are generated in memory order +(dimension 0 fastest), and a *block* is a run of that many consecutive elements, +so for a `(within, block)` layout blocks line up with dimension 0. + +| Constructor | Values | +| ------------------------------------------------ | -------------------------------------------------------------------------------------------------------------- | +| `uniform(lo, hi)`, `uniform_int(lo, hi)` | Uniform floats in `[lo, hi)` / integers in `[lo, hi]`. | +| `normal(mean, stddev)`, `constant(v)`, `zeros()` | Approximately normal (support `mean +- 6 stddev`); constant. | +| `blockwise_constant(block, base)` | One value from `base` per block. | +| `outliers(base, block, magnitude)` | `base`, except one random element per block (random sign) is `magnitude`. | +| `extremes(type)` | Only the extreme values of `type`. | +| `special_floats()` | `+-0`, denormals, smallest normal, `+-inf`, NaN. Never part of another distribution. | +| `mixture(parts, block = 1)` | Each run of `block` elements drawn from one randomly chosen part. | +| `from_port(port, type)` | Uniform over the port's declared range; else the whole range of an integer type, or `normal(0, 1)` for floats. | + +**Round trips.** `verify_round_trip(a, inputs | input | dist, ...)` runs +`decode(encode(x))` and returns a `RoundTripReport`: value count, maximum +absolute and relative error (`|y - x| / max(|x|, floor)`), RMSE, exact-match +count, the worst coordinate and its values, whether a bound is declared and how +many values of the first input exceed it, and the seed and distribution as +provenance. + +**Properties.** A `Property` is a named check on a `PropertyContext` (the +inputs, the decoded values, the encoded values and ports, and a lazily computed +re-encoding). The library: + +| Property | Checks | +| ---------------------------------------- | ----------------------------------------------------------------------------------------- | +| `lossless()` | `decode(encode(x))` is bit-for-bit `x`. | +| `bounded_error(abs)` | `abs(decode(encode(x)) - x) <= abs`. | +| `within_declared_bound()` | Same, against the unit's `error_bound()`; fails if none is declared. | +| `idempotent_requantize(rel_tol = 1e-6)` | `encode(decode(e)) == e` for `e = encode(x)` (floating-point encodings within `rel_tol`). | +| `zero_preserving()`, `sign_preserving()` | Zero decodes to zero; the round trip never flips a sign. | +| `outputs_within_declared_ranges()` | Every encoded value lies in its port's declared output range. | + +`check_property(a, prop, ...)` runs `prop` over `trials` seeded trials (default +8), each on freshly generated inputs, and stops at the first failure. Inputs +come from a `std::vector{type, extents, dist}`, from a single +`(dist, type, extents)`, or, with only `extents`, from the root's declared input +ports via `Distribution::from_port()`, so the property is exercised exactly +where its preconditions hold. Trial `i` uses `derive_seed(seed, i)` with trial 0 +using `seed`, so `trials = 1, seed = failing_seed` reproduces a failure. The +`PropertyResult` carries the property, stage label, trials run, failing seed and +coordinate, and message. + +**Preconditions.** The values entering the checked unit are compared with its +declared input ranges. By default a trial with a value outside them is *not run* +and fails with `PropertyResult::precondition_violated` set: the property is only +claimed where its preconditions hold, so a generator (or upstream stage) that +breaks them is reported as such. `PropertyOptions::check_preconditions = false` +runs the property anyway, which is how to demonstrate that it fails outside +them. + +**Conditioning on upstream guarantees.** `prop.at(stage)` runs the whole scheme +on the generated inputs, takes the values that actually arrive at `stage`'s +encode inputs, and checks the property for that stage alone (encode, then +decode) on them. `stage` must be invoked exactly once by the scheme, so hold on +to its handle. This shows whether an upstream quantizer really establishes the +precondition of the packing stage below it: + +```cpp +check_ranges(scheme); // empty: the ranges line up statically +check_property(scheme, lossless().at(pack), Distribution::normal(0, 1), + Float(32), {64}).passed; // pack is exact on what it receives +``` + +**Mechanics.** Stage-boundary values are read back by tracing stores (a JIT +custom trace handler) rather than by realizing the Funcs, since their extents +are inferred from the consumers. Each trial is a full round trip; idempotence +re-runs it on the decoded values in a separate pipeline. Funcs that are not +single-valued scalars (Tuple-valued or struct-typed) cannot be read back: their +`encoded()` Buffers are undefined, and a property that needs one (e.g. +`idempotent_requantize`, or `outputs_within_declared_ranges` on a port with a +range) fails with a message saying so, as does `prop.at(stage)` for a stage with +such an input. Properties over the decoded values (`lossless()`, +`bounded_error()`, ...) are unaffected. + +## Summary of decisions and open items + +| Item | Status | +| -------------------------------------------------------------------------------------------- | ------------------------------------------------------------------------------- | +| `Approximation`: value-semantic, type-erased, duck-typed handle; operates on `Func`s only | Decided | +| Copies of a handle are one stage; separate conversions are distinct stages | Decided; lookups (`encoded_by`/`decoded_by`) are by handle | +| Units return only their outputs; intermediates are discovered by walking definitions | Decided; scheduling-only, kept separate from the signature-contract outputs | +| Three unit forms per direction (single, multi, port-aware) | Decided; see API surface below | +| Named ports, optional declared signatures (static or contextual), run-time validation | Decided | +| `decode(encode(f))` reproduces `f`'s arg list and type | Decided; checked only at `approximate_by`'s substitution point, not generically | +| `Approximation` makes no placement claims (offline vs fused) | Decided | +| `encode`'s output arity/layout (packed vs planar) | Left to each `Approximation` | +| `approximate_by`: eager, destructive `substitute_calls`, not `Func::in`; a `Func` member | Decided; same scoping as `rfactor`, explicit already-existing `consumers` | +| `sever`: seam exposure on `Pipeline`, adopted via `add_input(ImageParam)`/`add_output(Func)` | Decided; single-valued Funcs only | +| `sever`: true automatic pipeline splitting | Rejected (phase-ordering conflict with `configure()`/`generate()`) | +| `sever`: cross-compile provenance checking | Deferred; relies on both sides building the same scheme | +| Fusing `encode`/`decode` into neighboring stages (activation requantization) | No new mechanism; ordinary `.compute_at()`/`.compute_inline()` | +| Declared ranges, `error_bound()`, `lossless()`; property-based testing in `tools/` | Decided; claims are checked in tests, never enforced or used in codegen | +| Generator I/O ergonomics (`add_input`/`add_output` pointer bookkeeping) | Accepted rough edge, deferred | + +Open items: + +- **API surface.** The unit interface has many optional hooks. Candidates for + pruning if the surface proves too large: the port-aware `encode`/`decode` + forms (needed only by combinators that route by name or forward context), + contextual `signature(inputs)` (needed only by shape-polymorphic units and + combinators), and `child_inputs()` (only refines + `describe()`/`check_ranges()`). +- **Path-based stage lookup.** Stages are found only by handle. Selecting a + stage by its path in the trace tree (e.g. by label chain) is not implemented; + it would remove the need to name handles in the common case, at the cost of + making lookups depend on labels. +- **Trace scope.** The trace collector is thread-local; handle calls made from + another thread inside a unit are not attached to the enclosing call. +- **Bounds composition.** `Compose` declares no `error_bound` for lossy stages, + and ranges flow only through units that declare them. +- **Provenance checking** across an offline/online split, as above. + +## Prior art referenced + +- An out-of-tree GGML application's hand-written quantize/dequantize/vec_dot + Generators, the implementations this design generalizes, and the first user of + this API. +- A private Python research prototype exploring the same compositional + `Approximation` idea. It is not a public artifact and is referenced only for + context. +- `apps/hannk/halide/conv_generator.cpp`: the in-repo precedent for a + `configure()` that does more than tweak a type. +- `src/Func.cpp: Stage::rfactor`: the eager, destructive graph-editing precedent + `approximate_by` follows instead of `.in()`. +- `src/Func.cpp`, `src/Function.cpp`, `src/WrapCalls.cpp`: origin of + `Function::substitute_calls`, the primitive `approximate_by` calls directly + and eagerly instead of through the deferred wrapper map. +- `src/Generator.h`: the `configure()`/`generate()`/`schedule()` lifecycle that + `sever` and the Generator shape build on. diff --git a/doc/CMakeLists.txt b/doc/CMakeLists.txt index 13ba994ac997..4bc445a58598 100644 --- a/doc/CMakeLists.txt +++ b/doc/CMakeLists.txt @@ -121,6 +121,7 @@ set(_sphinx_static_files "${CMAKE_CURRENT_SOURCE_DIR}/conf.py" "${CMAKE_CURRENT_SOURCE_DIR}/index.md" "${CMAKE_CURRENT_SOURCE_DIR}/guides.md" + "${CMAKE_CURRENT_SOURCE_DIR}/Approximation.md" "${CMAKE_CURRENT_SOURCE_DIR}/BuildingHalideWithCMake.md" "${CMAKE_CURRENT_SOURCE_DIR}/CodeStyleCMake.md" "${CMAKE_CURRENT_SOURCE_DIR}/CustomRuntimes.md" diff --git a/doc/guides.md b/doc/guides.md index e483e17c4cb2..1bd0a22ce69d 100644 --- a/doc/guides.md +++ b/doc/guides.md @@ -7,6 +7,7 @@ BuildingHalideWithCMake HalideCMakePackage CodeStyleCMake +Approximation CustomRuntimes FuzzTesting GeneratorCache diff --git a/src/Approximation.cpp b/src/Approximation.cpp new file mode 100644 index 000000000000..746e21b6dbe1 --- /dev/null +++ b/src/Approximation.cpp @@ -0,0 +1,1408 @@ +#include "Approximation.h" + +#include +#include +#include + +#if defined(__GNUC__) || defined(__clang__) +#include +#endif + +#include "Error.h" +#include "FindCalls.h" +#include "Function.h" + +namespace Halide { + +namespace { + +const ApproximationStageOutputs &find_stage(const std::vector &outputs, + const Approximation &stage, const char *direction) { + user_assert(stage.defined()) << "Approximation::" << direction << "_by: undefined stage\n"; + const ApproximationStageOutputs *found = nullptr; + size_t count = 0; + for (const ApproximationStageOutputs &o : outputs) { + if (o.stage.same_as(stage)) { + found = &o; + count++; + } + } + user_assert(count > 0) + << "Approximation::" << direction << "_by: the stage was not invoked in the " + << direction << " direction (was it converted separately from the handle passed " + << "into the Approximation?)\n"; + user_assert(count == 1) + << "Approximation::" << direction << "_by: ambiguous: stage invoked " << count + << " times in the " << direction << " direction\n"; + return *found; +} + +Func find_stage_output(const std::vector &outputs, + const Approximation &stage, size_t port, const char *direction) { + const ApproximationStageOutputs &found = find_stage(outputs, stage, direction); + return port < found.ports.size() ? found.ports[port] : Func(); +} + +Func find_stage_output(const std::vector &outputs, + const Approximation &stage, const std::string &port, const char *direction) { + const ApproximationStageOutputs &found = find_stage(outputs, stage, direction); + std::string available; + size_t match = 0, count = 0; + for (size_t i = 0; i < found.port_names.size(); i++) { + available += (i ? ", " : "") + found.port_names[i]; + if (found.port_names[i] == port) { + match = i; + count++; + } + } + user_assert(count > 0) + << "Approximation::" << direction << "_by: stage '" << stage.label() << "' has no " + << direction << " port named '" << port << "' (available: " << available << ")\n"; + user_assert(count == 1) + << "Approximation::" << direction << "_by: stage '" << stage.label() << "' has " << count + << " " << direction << " ports named '" << port << "'\n"; + return found.ports[match]; +} + +} // namespace + +Func ApproximationResult::encoded_by(const Approximation &stage, size_t port) const { + return find_stage_output(encoded_stage_outputs, stage, port, "encode"); +} + +Func ApproximationResult::decoded_by(const Approximation &stage, size_t port) const { + return find_stage_output(decoded_stage_outputs, stage, port, "decode"); +} + +Func ApproximationResult::encoded_by(const Approximation &stage, const std::string &port) const { + return find_stage_output(encoded_stage_outputs, stage, port, "encode"); +} + +Func ApproximationResult::decoded_by(const Approximation &stage, const std::string &port) const { + return find_stage_output(decoded_stage_outputs, stage, port, "decode"); +} + +std::vector ApproximationResult::stage_ports() const { + std::vector result; + std::set seen; + if (replacement.defined()) { + seen.insert(replacement.name()); + } + for (const auto *outputs : {&encoded_stage_outputs, &decoded_stage_outputs}) { + for (const ApproximationStageOutputs &o : *outputs) { + for (const Func &p : o.ports) { + if (p.defined() && seen.insert(p.name()).second) { + result.push_back(p); + } + } + } + } + return result; +} + +bool ApproximationResult::is_stage_port(const Func &f) const { + if (!f.defined()) { + return false; + } + for (const Func &p : stage_ports()) { + if (p.name() == f.name()) { + return true; + } + } + return false; +} + +namespace { + +void replace_all(std::string &s, const std::string &from, const std::string &to) { + for (size_t pos = s.find(from); pos != std::string::npos; pos = s.find(from, pos + to.size())) { + s.replace(pos, from.size(), to); + } +} + +void print_names(std::ostream &stream, const std::vector &funcs, + const std::vector &port_names = {}) { + const char *sep = ""; + for (size_t i = 0; i < funcs.size(); i++) { + stream << sep; + if (i < port_names.size()) { + stream << port_names[i] << "="; + } + stream << (funcs[i].defined() ? funcs[i].name() : ""); + sep = ", "; + } +} + +void print_node(std::ostream &stream, const ApproximationTraceNode &node, int depth) { + std::string indent(depth * 2, ' '); + stream << indent << node.label << " -> "; + print_names(stream, node.ports, node.port_names); + stream << "\n"; + if (!node.intermediates.empty()) { + stream << indent << " intermediates: "; + print_names(stream, node.intermediates); + stream << "\n"; + } + for (const ApproximationTraceNode &child : node.children) { + print_node(stream, child, depth + 1); + } +} + +void flatten(const ApproximationTraceNode &node, std::vector &out) { + for (const ApproximationTraceNode &child : node.children) { + flatten(child, out); + } + out.push_back({node.stage, node.ports, node.port_names, node.intermediates}); +} + +} // namespace + +std::ostream &operator<<(std::ostream &stream, const ApproximationTraceNode &node) { + print_node(stream, node, 0); + return stream; +} + +std::ostream &operator<<(std::ostream &stream, const ApproximationResult &result) { + stream << "encode:\n"; + if (result.encode_trace.stage.defined()) { + print_node(stream, result.encode_trace, 1); + } + stream << "decode:\n"; + if (result.decode_trace.stage.defined()) { + print_node(stream, result.decode_trace, 1); + } + return stream; +} + +std::string Approximation::type_label(const char *pretty_function) { + const std::string pretty = pretty_function; + std::string name; + size_t start = std::string::npos; + if (size_t pos = pretty.find("pretty_type_name<"); pos != std::string::npos) { + // MSVC: ...pretty_type_name(void) + start = pos + std::string("pretty_type_name<").size(); + } else if (pos = pretty.find("T = "); pos != std::string::npos) { + // clang: [T = Foo]; GCC: [with T = Foo; ...] + start = pos + 4; + } + if (start == std::string::npos) { + return "unit"; + } + int depth = 0; + for (size_t i = start; i < pretty.size(); i++) { + char c = pretty[i]; + if (c == '<' || c == '(' || c == '[') { + depth++; + } else if (c == '>' || c == ')' || c == ']') { + if (depth == 0) { + break; + } + depth--; + } else if (c == ';' && depth == 0) { + break; + } + name += c; + } + replace_all(name, "struct ", ""); + replace_all(name, "class ", ""); + replace_all(name, "(anonymous namespace)::", ""); + replace_all(name, "{anonymous}::", ""); + replace_all(name, "`anonymous namespace'::", ""); + replace_all(name, "`anonymous-namespace'::", ""); + replace_all(name, "std::__1::", "std::"); + replace_all(name, "std::__cxx11::", "std::"); + replace_all(name, "Halide::", ""); + return name; +} + +Approximation Approximation::labelled(std::string label) const { + user_assert(defined()) << "labelled called on an undefined Approximation\n"; + state_->label = std::move(label); + return *this; +} + +std::string Approximation::label() const { + if (!defined()) { + return ""; + } + return state_->label.empty() ? state_->impl->default_label() : state_->label; +} + +void Approximation::check_single_input(const std::vector &inputs, const char *direction) { + user_assert(inputs.size() == 1) + << "Approximation: a unit with a single-Func " << direction << "() was given " + << inputs.size() << " inputs, but requires exactly one\n"; +} + +namespace { + +// The children collected so far by the innermost handle call in progress on +// this thread (null outside any call). Encode and decode share it. +thread_local std::vector *active_children = nullptr; + +struct CallScope { + std::vector children; + std::vector *parent; + + CallScope() + : parent(active_children) { + active_children = &children; + } + + ~CallScope() { + active_children = parent; + } + + // Finish the call: nest the node under the enclosing call, if any. + ApproximationTraceNode finish(const Approximation &stage, std::vector ports, + std::vector port_names, std::vector intermediates, + std::vector inputs, std::vector input_names) { + ApproximationTraceNode node{stage, stage.label(), std::move(ports), std::move(port_names), + std::move(intermediates), std::move(children), std::move(inputs), + std::move(input_names)}; + if (parent) { + parent->push_back(node); + } + return node; + } +}; + +// Post-order DFS over direct calls (visited by name order) so producers come +// before consumers. Inputs are never entered. +void discover(const Internal::Function &f, const std::set &inputs, + const std::set &outputs, std::set &visited, + std::vector &result) { + if (inputs.count(f.name()) || !visited.insert(f.name()).second) { + return; + } + for (const auto &[name, callee] : Internal::find_direct_calls(f)) { + discover(callee, inputs, outputs, visited, result); + } + if (!outputs.count(f.name())) { + result.emplace_back(f); + } +} + +std::vector find_intermediates(const std::vector &inputs, const std::vector &outputs) { + std::set input_names, output_names, visited; + for (const Func &f : inputs) { + if (f.defined()) { + input_names.insert(f.name()); + } + } + for (const Func &f : outputs) { + if (f.defined()) { + output_names.insert(f.name()); + } + } + std::vector result; + for (const Func &f : outputs) { + if (f.defined()) { + discover(f.function(), input_names, output_names, visited, result); + } + } + return result; +} + +} // namespace + +namespace { + +std::string range_string(const ApproximationRange &range) { + std::ostringstream stream; + stream.precision(9); + stream << "[" << range.lo << ", " << range.hi << "]"; + return stream.str(); +} + +std::string port_string(const ApproximationPort &port) { + std::ostringstream stream; + stream << port.name; + if (port.type) { + stream << ": " << *port.type; + } + if (port.dimensions) { + stream << " x" << *port.dimensions; + } + if (port.range) { + stream << " in " << range_string(*port.range); + } + return stream.str(); +} + +// Compare the ranges `context` guarantees with the preconditions `required` +// declares, port by port. Nothing is reported unless the two line up. +std::vector input_range_issues(const ApproximationPorts &required, + const ApproximationPorts &context) { + std::vector issues; + if (required.size() != context.size()) { + return issues; + } + for (size_t i = 0; i < required.size(); i++) { + if (!required[i].range) { + continue; + } + const std::string prefix = "input '" + required[i].name + "': "; + if (!context[i].range) { + issues.push_back(prefix + "requires " + range_string(*required[i].range) + + ", but the producer declares no range"); + } else if (!required[i].range->contains(*context[i].range)) { + issues.push_back(prefix + range_string(*context[i].range) + " not within " + + range_string(*required[i].range)); + } + } + return issues; +} + +std::string type_string(const Type &t) { + std::ostringstream stream; + stream << t; + return stream.str(); +} + +std::string ports_string(const ApproximationPorts &ports) { + std::string result; + for (size_t i = 0; i < ports.size(); i++) { + result += (i ? ", " : "") + port_string(ports[i]); + } + return result; +} + +ApproximationPorts positional_ports(size_t count) { + ApproximationPorts ports; + for (size_t i = 0; i < count; i++) { + ports.emplace_back(std::to_string(i)); + } + return ports; +} + +ApproximationPorts names_only(const ApproximationPorts &ports) { + ApproximationPorts result; + for (const ApproximationPort &p : ports) { + result.emplace_back(p.name); + } + return result; +} + +std::vector names_of(const ApproximationPorts &ports) { + std::vector names; + for (const ApproximationPort &p : ports) { + names.push_back(p.name); + } + return names; +} + +// Fill in whatever the ports leave unset from the actual Funcs. +void fill_from_funcs(ApproximationPorts &ports, const std::vector &funcs) { + for (size_t i = 0; i < ports.size(); i++) { + if (!funcs[i].defined()) { + continue; + } + if (!ports[i].type && funcs[i].outputs() == 1) { + ports[i].type = funcs[i].types()[0]; + } + if (!ports[i].dimensions) { + ports[i].dimensions = funcs[i].dimensions(); + } + } +} + +void validate_ports(const Approximation &stage, const char *direction, const char *role, + const std::vector &funcs, const ApproximationPorts &ports) { + for (size_t i = 0; i < ports.size(); i++) { + const Func &f = funcs[i]; + if (!f.defined()) { + continue; + } + const ApproximationPort &p = ports[i]; + if (p.type) { + user_assert(f.outputs() == 1 && f.types()[0] == *p.type) + << "Approximation '" << stage.label() << "' " << direction << ": " << role + << " port '" << p.name << "' expects type " << *p.type << " but Func '" << f.name() + << "' has " + << (f.outputs() == 1 ? "type " + type_string(f.types()[0]) : + "a Tuple of " + std::to_string(f.outputs()) + " values") + << "\n"; + } + if (p.dimensions) { + user_assert(f.dimensions() == *p.dimensions) + << "Approximation '" << stage.label() << "' " << direction << ": " << role + << " port '" << p.name << "' expects " << *p.dimensions << " dimensions but Func '" + << f.name() << "' has " << f.dimensions() << "\n"; + } + } +} + +} // namespace + +ApproximationPorts Approximation::resolve_inputs(const std::vector &inputs, const ApproximationPorts &given) const { + const size_t n = inputs.size(); + if (!given.empty()) { + user_assert(given.size() == n) + << "Approximation '" << label() << "' encode: " << given.size() + << " ports were given for " << n << " Funcs\n"; + validate_ports(*this, "encode", "input", inputs, given); + } + + // The declared inputs are defaults for the names, and constrain the rest. + ApproximationSignature sig = signature(given); + const bool have_declared = sig.known && state_->impl->signature_form() != SignatureForm::None; + if (have_declared) { + user_assert(sig.inputs.size() == n) + << "Approximation '" << label() << "' encode: the declared signature has " + << sig.inputs.size() << " input ports (" << ports_string(sig.inputs) << ") but " << n + << " Funcs were given\n"; + } + + ApproximationPorts ports = !given.empty() ? given : have_declared ? sig.inputs : + positional_ports(n); + if (have_declared && !given.empty()) { + for (size_t i = 0; i < n; i++) { + ports[i].type = sig.inputs[i].type ? sig.inputs[i].type : ports[i].type; + ports[i].dimensions = sig.inputs[i].dimensions ? sig.inputs[i].dimensions : ports[i].dimensions; + } + } + validate_ports(*this, "encode", "input", inputs, ports); + fill_from_funcs(ports, inputs); + return ports; +} + +ApproximationPorts Approximation::resolve_encoded(const std::vector &encoded, const ApproximationSignature &sig, + const ApproximationPorts &input_ports) const { + const size_t n = encoded.size(); + ApproximationPorts ports; + if (sig.known) { + user_assert(sig.outputs.size() == n) + << "Approximation '" << label() << "' decode: the declared signature has " + << sig.outputs.size() << " output ports (" << ports_string(sig.outputs) << ") but " << n + << " encoded Funcs were given\n"; + ports = sig.outputs; + } else { + ports = n == input_ports.size() ? names_only(input_ports) : positional_ports(n); + } + validate_ports(*this, "decode", "encoded", encoded, ports); + fill_from_funcs(ports, encoded); + return ports; +} + +ApproximationPorts Approximation::output_ports(const std::vector &outputs, const ApproximationSignature &sig, + const ApproximationPorts &input_ports, bool encode_direction) const { + const char *direction = encode_direction ? "encode" : "decode"; + ApproximationPorts ports; + if (sig.known && state_->impl->signature_form() != SignatureForm::None) { + ports = encode_direction ? sig.outputs : sig.inputs; + user_assert(ports.size() == outputs.size()) + << "Approximation '" << label() << "' " << direction << ": the declared signature has " + << ports.size() << (encode_direction ? " output" : " input") << " ports (" << ports_string(ports) + << ") but the unit returned " << outputs.size() << " Funcs\n"; + } else if (encode_direction) { + ports = outputs.size() == input_ports.size() ? names_only(input_ports) : positional_ports(outputs.size()); + } else if (sig.known) { + // Undeclared single-Func unit: its one output is named like its input. + ports = names_only(sig.inputs); + } else { + ports = outputs.size() == input_ports.size() ? names_only(input_ports) : positional_ports(outputs.size()); + } + validate_ports(*this, direction, "output", outputs, ports); + fill_from_funcs(ports, outputs); + return ports; +} + +ApproximationSignature Approximation::signature(const ApproximationPorts &inputs) const { + user_assert(defined()) << "signature called on an undefined Approximation\n"; + ApproximationSignature result; + switch (state_->impl->signature_form()) { + case SignatureForm::Static: + result = state_->impl->declared_signature({}); + break; + case SignatureForm::Contextual: + result = state_->impl->declared_signature(inputs); + break; + case SignatureForm::None: + if (!state_->impl->encode_is_single() || inputs.size() > 1) { + return ApproximationSignature::unknown(inputs); + } + result.inputs = inputs.empty() ? positional_ports(1) : inputs; + result.outputs = {ApproximationPort(result.inputs[0].name)}; + return result; + } + // Names that flow in win over declared ones. + if (result.known && !inputs.empty() && result.inputs.size() == inputs.size()) { + for (size_t i = 0; i < inputs.size(); i++) { + result.inputs[i].name = inputs[i].name; + } + } + return result; +} + +void Approximation::describe_to(std::string &out, const ApproximationPorts &inputs, int depth) const { + out += std::string(depth * 2, ' ') + label() + " "; + ApproximationSignature sig = signature(inputs); + if (sig.known) { + out += "(" + ports_string(sig.inputs) + ") -> (" + ports_string(sig.outputs) + ")"; + } else { + out += "(unknown signature)"; + } + out += '\n'; + if (sig.known) { + for (const std::string &issue : input_range_issues(sig.inputs, inputs)) { + out += std::string(depth * 2 + 2, ' '); + out += "! " + issue + "\n"; + } + } + std::vector children = state_->impl->children(); + std::vector contexts = state_->impl->child_inputs(inputs); + for (size_t i = 0; i < children.size(); i++) { + if (children[i].defined()) { + children[i].describe_to(out, i < contexts.size() ? contexts[i] : ApproximationPorts{}, depth + 1); + } + } +} + +void Approximation::range_issues_to(std::vector &out, const ApproximationPorts &inputs, + const std::string &path) const { + std::string here = path; + if (!here.empty()) { + here += " > "; + } + here += label(); + ApproximationSignature sig = signature(inputs); + if (sig.known) { + for (const std::string &issue : input_range_issues(sig.inputs, inputs)) { + std::string line = here; + line += ": "; + line += issue; + out.push_back(std::move(line)); + } + } + std::vector children = state_->impl->children(); + std::vector contexts = state_->impl->child_inputs(inputs); + for (size_t i = 0; i < children.size(); i++) { + if (children[i].defined()) { + children[i].range_issues_to(out, i < contexts.size() ? contexts[i] : ApproximationPorts{}, here); + } + } +} + +std::vector check_ranges(const Approximation &a, const ApproximationPorts &inputs) { + std::vector out; + if (a.defined()) { + a.range_issues_to(out, inputs, ""); + } + return out; +} + +Func Approximation::error_bound(const std::vector &inputs, const std::vector &encoded) const { + user_assert(defined()) << "error_bound called on an undefined Approximation\n"; + Func bound = state_->impl->error_bound(inputs, encoded); + if (bound.defined()) { + return bound; + } + if (!state_->impl->lossless() || inputs.empty() || !inputs[0].defined()) { + return Func(); + } + std::vector args; + args.reserve(inputs[0].dimensions()); + for (int i = 0; i < inputs[0].dimensions(); i++) { + args.emplace_back("zb" + std::to_string(i)); + } + Func zero("approximation_zero_bound"); + zero(args) = cast(0); + return zero; +} + +bool Approximation::lossless() const { + user_assert(defined()) << "lossless called on an undefined Approximation\n"; + return state_->impl->lossless(); +} + +std::string Approximation::describe(const ApproximationPorts &inputs) const { + if (!defined()) { + return "\n"; + } + std::string out; + describe_to(out, inputs, 0); + return out; +} + +std::ostream &operator<<(std::ostream &stream, const Approximation &approximation) { + return stream << approximation.describe(); +} + +EncodeResult Approximation::encode(const std::vector &inputs, const ApproximationPorts &input_ports) const { + user_assert(defined()) << "encode called on an undefined Approximation\n"; + ApproximationPorts resolved = resolve_inputs(inputs, input_ports); + CallScope scope; + std::vector encoded = state_->impl->encode(inputs, resolved); + ApproximationPorts encoded_ports = output_ports(encoded, signature(resolved), resolved, true); + std::vector intermediates = find_intermediates(inputs, encoded); + ApproximationTraceNode node = scope.finish(*this, encoded, names_of(encoded_ports), intermediates, inputs, names_of(resolved)); + std::vector stage_outputs; + flatten(node, stage_outputs); + return {std::move(encoded), std::move(encoded_ports), std::move(intermediates), std::move(stage_outputs), + std::move(node)}; +} + +DecodeResult Approximation::decode(const std::vector &encoded, const ApproximationPorts &input_ports) const { + user_assert(defined()) << "decode called on an undefined Approximation\n"; + // The context is static: the caller's, or else the defaults. + ApproximationPorts context = input_ports; + if (context.empty()) { + ApproximationSignature defaults = signature(); + context = defaults.known ? defaults.inputs : ApproximationPorts{}; + } + ApproximationSignature sig = signature(context); + ApproximationPorts resolved = resolve_encoded(encoded, sig, context); + CallScope scope; + std::vector decoded = state_->impl->decode(encoded, context); + ApproximationPorts decoded_ports = output_ports(decoded, sig, context, false); + std::vector intermediates = find_intermediates(encoded, decoded); + ApproximationTraceNode node = scope.finish(*this, decoded, names_of(decoded_ports), intermediates, encoded, names_of(resolved)); + std::vector stage_outputs; + flatten(node, stage_outputs); + return {std::move(decoded), std::move(decoded_ports), std::move(intermediates), std::move(stage_outputs), + std::move(node)}; +} + +std::vector Compose::encode(const std::vector &inputs, const ApproximationPorts &input_ports) const { + user_assert(!stages.empty()) << "Compose::encode: no stages\n"; + std::vector current = inputs; + ApproximationPorts current_ports = input_ports; + for (const Approximation &stage : stages) { + EncodeResult r = stage.encode(current, current_ports); + current = std::move(r.encoded); + current_ports = std::move(r.encoded_ports); + } + return current; +} + +std::vector Compose::decode(const std::vector &encoded, const ApproximationPorts &input_ports) const { + user_assert(!stages.empty()) << "Compose::decode: no stages\n"; + std::vector contexts = child_inputs(input_ports); + std::vector current = encoded; + for (int i = (int)stages.size() - 1; i >= 0; i--) { + current = stages[i].decode(current, contexts[i]).decoded; + } + return current; +} + +ApproximationSignature Compose::signature(const ApproximationPorts &inputs) const { + if (stages.empty()) { + return ApproximationSignature::unknown(inputs); + } + ApproximationPorts current = inputs, first_inputs; + for (size_t i = 0; i < stages.size(); i++) { + ApproximationSignature s = stages[i].signature(current); + if (!s.known) { + return ApproximationSignature::unknown(inputs); + } + if (i == 0) { + first_inputs = s.inputs; + } + current = std::move(s.outputs); + } + ApproximationSignature result; + result.inputs = std::move(first_inputs); + result.outputs = std::move(current); + return result; +} + +std::vector Compose::children() const { + return stages; +} + +std::vector Compose::child_inputs(const ApproximationPorts &inputs) const { + std::vector result; + ApproximationPorts current = inputs; + for (const Approximation &stage : stages) { + result.push_back(current); + ApproximationSignature s = stage.signature(current); + current = s.known ? std::move(s.outputs) : ApproximationPorts{}; + } + return result; +} + +bool Compose::lossless() const { + for (const Approximation &stage : stages) { + if (!stage.lossless()) { + return false; + } + } + return !stages.empty(); +} + +bool Choose::lossless() const { + return chosen.lossless(); +} + +namespace { + +template +std::vector slice(const std::vector &v, size_t begin, size_t count) { + return std::vector(v.begin() + begin, v.begin() + begin + count); +} + +// A port passed through a by-name Parallel is not a precondition of the +// Parallel: keep the interface, drop the range. +ApproximationPort without_range(ApproximationPort port) { + port.range.reset(); + return port; +} + +// The number of inputs `child` takes when it has no context (one if unknown). +size_t input_width(const Approximation &child) { + ApproximationSignature s = child.signature(); + return s.known && !s.inputs.empty() ? s.inputs.size() : 1; +} + +// The number of encoded ports `child` produces from `context` (one if unknown). +size_t output_width(const Approximation &child, const ApproximationPorts &context) { + ApproximationSignature s = child.signature(context); + return s.known && !s.outputs.empty() ? s.outputs.size() : 1; +} + +} // namespace + +Parallel::Parallel(std::initializer_list entries) { + for (const Entry &e : entries) { + ports_.push_back(e.port); + children_.push_back(e.child); + } +} + +bool Parallel::plan(const ApproximationPorts &inputs, std::vector &segments, std::string *problem) const { + auto fail = [&](const std::string &message) { + if (problem) { + *problem = message; + } + return false; + }; + segments.clear(); + if (ports_.empty()) { + size_t begin = 0; + for (size_t i = 0; i < children_.size(); i++) { + const size_t width = input_width(children_[i]); + segments.push_back({(int)i, begin, width, {}, 0}); + begin += width; + } + if (!inputs.empty() && begin != inputs.size()) { + return fail("the children take " + std::to_string(begin) + " Funcs in total, but there are " + + std::to_string(inputs.size()) + " (the ports are: " + ports_string(names_only(inputs)) + ")"); + } + for (Segment &seg : segments) { + if (!inputs.empty()) { + seg.context = slice(inputs, seg.begin, seg.width); + } + seg.outputs = output_width(children_[seg.child], seg.context); + } + return true; + } + + std::vector owner(inputs.size(), -1); + for (size_t i = 0; i < ports_.size(); i++) { + size_t match = 0, count = 0; + for (size_t j = 0; j < inputs.size(); j++) { + if (inputs[j].name == ports_[i]) { + match = j; + count++; + } + } + if (count != 1) { + return fail(std::string("there is ") + (count == 0 ? "no" : "more than one") + " input port named '" + + ports_[i] + "' (the ports are: " + ports_string(names_only(inputs)) + ")"); + } + if (owner[match] >= 0) { + return fail("the port '" + ports_[i] + "' is routed by more than one entry"); + } + owner[match] = (int)i; + } + for (size_t j = 0; j < inputs.size(); j++) { + Segment seg{owner[j], j, 1, {inputs[j]}, 1}; + if (owner[j] >= 0) { + seg.outputs = output_width(children_[owner[j]], seg.context); + } + segments.push_back(seg); + } + return true; +} + +std::vector Parallel::encode(const std::vector &inputs, const ApproximationPorts &input_ports) const { + ApproximationPorts ports = input_ports.size() == inputs.size() ? input_ports : positional_ports(inputs.size()); + std::vector segments; + std::string problem; + user_assert(plan(ports, segments, &problem)) << "Parallel::encode: " << problem << "\n"; + std::vector result; + for (const Segment &seg : segments) { + std::vector in = slice(inputs, seg.begin, seg.width); + if (seg.child < 0) { + result.insert(result.end(), in.begin(), in.end()); + } else { + std::vector out = children_[seg.child].encode(in, seg.context).encoded; + result.insert(result.end(), out.begin(), out.end()); + } + } + return result; +} + +std::vector Parallel::decode(const std::vector &encoded, const ApproximationPorts &input_ports) const { + std::vector segments; + std::string problem; + user_assert(plan(input_ports, segments, &problem)) << "Parallel::decode: " << problem << "\n"; + size_t total = 0; + for (const Segment &seg : segments) { + total += seg.outputs; + } + user_assert(total == encoded.size()) + << "Parallel::decode: the children produce " << total << " Funcs in total, but " << encoded.size() + << " were given\n"; + std::vector result; + size_t begin = 0; + for (const Segment &seg : segments) { + std::vector in = slice(encoded, begin, seg.outputs); + begin += seg.outputs; + if (seg.child < 0) { + result.insert(result.end(), in.begin(), in.end()); + continue; + } + std::vector out = children_[seg.child].decode(in, seg.context).decoded; + user_assert(ports_.empty() || out.size() == 1) + << "Parallel::decode: the child for port '" << ports_[seg.child] << "' decoded to " << out.size() + << " Funcs, but must decode to one\n"; + result.insert(result.end(), out.begin(), out.end()); + } + return result; +} + +ApproximationSignature Parallel::signature(const ApproximationPorts &inputs) const { + std::vector segments; + if (children_.empty() || (!ports_.empty() && inputs.empty()) || !plan(inputs, segments, nullptr)) { + return ApproximationSignature::unknown(inputs); + } + ApproximationSignature result; + for (const Segment &seg : segments) { + if (seg.child < 0) { + result.inputs.push_back(without_range(seg.context[0])); + result.outputs.push_back(seg.context[0]); + continue; + } + ApproximationSignature s = children_[seg.child].signature(seg.context); + if (!s.known || s.inputs.size() != seg.width) { + return ApproximationSignature::unknown(inputs); + } + result.inputs.insert(result.inputs.end(), s.inputs.begin(), s.inputs.end()); + result.outputs.insert(result.outputs.end(), s.outputs.begin(), s.outputs.end()); + } + return result; +} + +std::vector Parallel::children() const { + return children_; +} + +std::vector Parallel::child_inputs(const ApproximationPorts &inputs) const { + std::vector result(children_.size()); + std::vector segments; + if (plan(inputs, segments, nullptr)) { + for (const Segment &seg : segments) { + if (seg.child >= 0) { + result[seg.child] = seg.context; + } + } + } + return result; +} + +bool Parallel::lossless() const { + for (const Approximation &child : children_) { + if (!child.lossless()) { + return false; + } + } + return !children_.empty(); +} + +std::vector TrustedInverse::encode(const std::vector &inputs, const ApproximationPorts &input_ports) const { + return encoder.encode(inputs, input_ports).encoded; +} + +std::vector TrustedInverse::decode(const std::vector &encoded, const ApproximationPorts &input_ports) const { + return decoder.decode(encoded, input_ports).decoded; +} + +ApproximationSignature TrustedInverse::signature(const ApproximationPorts &inputs) const { + return encoder.signature(inputs); +} + +std::vector TrustedInverse::children() const { + return {encoder, decoder}; +} + +std::vector TrustedInverse::child_inputs(const ApproximationPorts &inputs) const { + return {inputs, {}}; +} + +std::vector Choose::encode(const std::vector &inputs, const ApproximationPorts &input_ports) const { + return chosen.encode(inputs, input_ports).encoded; +} + +std::vector Choose::decode(const std::vector &encoded, const ApproximationPorts &input_ports) const { + return chosen.decode(encoded, input_ports).decoded; +} + +ApproximationSignature Choose::signature(const ApproximationPorts &inputs) const { + return chosen.signature(inputs); +} + +std::vector Choose::children() const { + return {chosen}; +} + +std::vector Choose::child_inputs(const ApproximationPorts &inputs) const { + return {inputs}; +} + +namespace { + +Pointwise::TupleFn wrap_expr_fn(Pointwise::ExprFn fn) { + return [fn = std::move(fn)](const std::vector &values) { + user_assert(values.size() == 1) + << "Pointwise: an Expr-valued function requires a single-output Func, but the input has " + << values.size() << " outputs (use a std::vector function instead)\n"; + return std::vector{fn(values[0])}; + }; +} + +Func apply_pointwise(const Func &input, const Pointwise::TupleFn &fn, const std::string &name, + const std::string &var_prefix, const char *direction) { + user_assert(input.defined()) << "Pointwise::" << direction << ": undefined Func\n"; + std::vector args; + std::vector call_args; + for (int i = 0; i < input.dimensions(); i++) { + args.emplace_back(var_prefix + std::to_string(i)); + call_args.emplace_back(args.back()); + } + std::vector values; + if (input.outputs() == 1) { + values.push_back(input(call_args)); + } else { + for (int i = 0; i < input.outputs(); i++) { + values.push_back(input(call_args)[i]); + } + } + std::vector results = fn(values); + user_assert(!results.empty()) << "Pointwise::" << direction << ": function returned no values\n"; + Func out(name); + if (results.size() == 1) { + out(args) = results[0]; + } else { + out(args) = Tuple(results); + } + return out; +} + +} // namespace + +Pointwise::Pointwise(const std::string &name, ExprFn encode_fn, ExprFn decode_fn) + : Pointwise(name + "_encode", name + "_decode", std::move(encode_fn), std::move(decode_fn)) { +} + +Pointwise::Pointwise(const std::string &name, TupleFn encode_fn, TupleFn decode_fn) + : Pointwise(name + "_encode", name + "_decode", std::move(encode_fn), std::move(decode_fn)) { +} + +Pointwise::Pointwise(std::string encode_name, std::string decode_name, ExprFn encode_fn, ExprFn decode_fn, + std::string var_prefix) + : Pointwise(std::move(encode_name), std::move(decode_name), + wrap_expr_fn(std::move(encode_fn)), wrap_expr_fn(std::move(decode_fn)), + std::move(var_prefix)) { +} + +Pointwise::Pointwise(std::string encode_name, std::string decode_name, TupleFn encode_fn, TupleFn decode_fn, + std::string var_prefix) + : encode_name(std::move(encode_name)), decode_name(std::move(decode_name)), + var_prefix(std::move(var_prefix)), encode_fn(std::move(encode_fn)), decode_fn(std::move(decode_fn)) { +} + +Func Pointwise::encode(const Func &input) const { + return apply_pointwise(input, encode_fn, encode_name, var_prefix, "encode"); +} + +Func Pointwise::decode(const Func &encoded) const { + return apply_pointwise(encoded, decode_fn, decode_name, var_prefix, "decode"); +} + +Pointwise Pointwise::with_error_bound(ExprFn bound) const { + Pointwise copy = *this; + copy.bound_fn = std::move(bound); + return copy; +} + +Pointwise Pointwise::with_types(Type input_type, Type output_type) const { + Pointwise copy = *this; + copy.input_type = input_type; + copy.output_type = output_type; + return copy; +} + +Pointwise Pointwise::with_ranges(ApproximationRange input_range, ApproximationRange output_range) const { + Pointwise copy = *this; + copy.input_range = input_range; + copy.output_range = output_range; + return copy; +} + +Pointwise Pointwise::with_lossless(bool is_lossless) const { + Pointwise copy = *this; + copy.lossless_ = is_lossless; + return copy; +} + +ApproximationSignature Pointwise::signature(const ApproximationPorts &inputs) const { + if (inputs.size() > 1) { + return ApproximationSignature::unknown(inputs); + } + std::optional dims; + std::string name = positional_ports(1)[0].name; + if (inputs.size() == 1) { + dims = inputs[0].dimensions; + name = inputs[0].name; + } + return {{{name, input_type, dims, input_range}}, {{name, output_type, dims, output_range}}}; +} + +Func Pointwise::error_bound(const std::vector &inputs, const std::vector &) const { + if (!bound_fn) { + return Func(); + } + user_assert(inputs.size() == 1 && inputs[0].outputs() == 1) + << "Pointwise::error_bound requires a single-valued input Func\n"; + return apply_pointwise( + inputs[0], wrap_expr_fn([f = bound_fn](const Expr &x) { return f(x); }), + encode_name + "_bound", var_prefix, "error_bound"); +} + +std::vector Identity::encode(const std::vector &inputs) const { + return inputs; +} + +std::vector Identity::decode(const std::vector &encoded) const { + return encoded; +} + +ApproximationSignature Identity::signature(const ApproximationPorts &inputs) const { + if (inputs.empty()) { + return ApproximationSignature::unknown(); + } + return {inputs, inputs}; +} + +ApproximationSignature Permute::signature(const ApproximationPorts &inputs) const { + ApproximationPorts ports = inputs.empty() ? positional_ports(forward.size()) : inputs; + if (ports.size() != forward.size()) { + return ApproximationSignature::unknown(inputs); + } + ApproximationSignature result; + result.inputs = ports; + for (int source : forward) { + result.outputs.push_back(ports[source]); + } + return result; +} + +std::vector Permute::encode(const std::vector &inputs) const { + user_assert(inputs.size() == forward.size()) << "Permutation size does not match input size"; + std::vector result; + result.reserve(inputs.size()); + for (int i = 0; i < (int)inputs.size(); i++) { + result.push_back(inputs[forward[i]]); + } + return result; +} + +std::vector Permute::decode(const std::vector &encoded) const { + user_assert(encoded.size() == forward.size()) << "Permutation size does not match encoded size"; + std::vector result; + result.reserve(encoded.size()); + for (int i = 0; i < (int)encoded.size(); i++) { + result.push_back(encoded[backward[i]]); + } + return result; +} + +namespace { + +std::vector component_vars(int dimensions, const std::string &prefix) { + std::vector vars; + vars.reserve(dimensions); + for (int i = 0; i < dimensions; ++i) { + vars.emplace_back(prefix + std::to_string(i)); + } + return vars; +} + +std::vector component_exprs(const std::vector &vars) { + return std::vector(vars.begin(), vars.end()); +} + +} // namespace + +std::vector BlockReshape::encode(const std::vector &inputs) const { + user_assert(inputs.size() == 1) << "BlockReshape::encode expects one input\n"; + const Func &flat = inputs[0]; + std::vector dims = block_vars(); + Var blk("blk"); + Expr within = cast(0); + int stride = 1; + for (size_t i = 0; i < dims.size(); ++i) { + within += dims[i] * stride; + stride *= extents_[i]; + } + std::vector args = dims; + args.push_back(blk); + Func packed("block_reshape_packed"); + packed(args) = block_indexed_ ? flat(within, blk) : flat(blk * block_size() + within); + return {packed}; +} + +std::vector BlockReshape::decode(const std::vector &encoded) const { + user_assert(encoded.size() == 1) << "BlockReshape::decode expects one input\n"; + const Func &packed = encoded[0]; + Var k("k"), kk("kk"), blk("blk"); + Expr within = block_indexed_ ? Expr(kk) : k % block_size(); + Expr block = block_indexed_ ? Expr(blk) : k / block_size(); + std::vector args; + Expr rem = within; + for (int extent : extents_) { + args.push_back(rem % extent); + rem /= extent; + } + args.push_back(block); + Func out("block_reshape_unpacked"); + if (block_indexed_) { + out(kk, blk) = packed(args); + } else { + out(k) = packed(args); + } + return {out}; +} + +ApproximationSignature BlockReshape::signature() const { + return {{{"values", std::nullopt, block_indexed_ ? 2 : 1}}, + {{"blocks", std::nullopt, (int)extents_.size() + 1}}}; +} + +int BlockReshape::block_size() const { + int size = 1; + for (int extent : extents_) { + size *= extent; + } + return size; +} + +std::vector BlockReshape::block_vars() const { + std::vector vars; + vars.reserve(extents_.size()); + for (size_t i = 0; i < extents_.size(); ++i) { + vars.emplace_back(extents_.size() == 1 ? "kk" : "d" + std::to_string(i)); + } + return vars; +} + +StructLayout::StructLayout(Type record_type, std::vector logical_fields, + int record_dimensions) + : record_type_(record_type), logical_fields_(std::move(logical_fields)), + record_dimensions_(record_dimensions) { + user_assert(record_type_.is_struct()) << "StructLayout requires a struct Type\n"; + user_assert(record_dimensions_ > 0) << "StructLayout record dimensionality must be positive\n"; + const StructTypeInfo *info = record_type_.struct_type(); + user_assert(logical_fields_.size() == info->fields.size()) + << "StructLayout requires exactly one logical slot per physical field\n"; + for (const std::string &name : logical_fields_) { + int matches = 0; + for (const StructField &field : info->fields) { + matches += field.name == name; + } + user_assert(matches == 1) << "StructLayout: no unique field named '" << name << "'\n"; + int logical_matches = 0; + for (const std::string &logical_name : logical_fields_) { + logical_matches += logical_name == name; + } + user_assert(logical_matches == 1) << "StructLayout: duplicate logical field '" << name << "'\n"; + } +} + +std::vector StructLayout::encode(const std::vector &inputs) const { + user_assert(inputs.size() == logical_fields_.size()) + << "StructLayout::encode input count does not match logical field count\n"; + const StructTypeInfo *info = record_type_.struct_type(); + std::vector records = component_vars(record_dimensions_, "record"); + std::vector record_args = component_exprs(records); + std::vector values; + for (const StructField &field : info->fields) { + size_t slot = logical_slot(field.name); + const Func &input = inputs[slot]; + user_assert(input.outputs() == 1 && input.types()[0] == field.type) + << "StructLayout field '" << field.name << "' requires exact type " << field.type + << " but slot " << slot << " has " << input.types()[0] << "\n"; + int extent = field.array_extent.value_or(1); + user_assert(input.dimensions() == record_dimensions_ + (field.array_extent ? 1 : 0)) + << "StructLayout field '" << field.name << "' has the wrong dimensionality\n"; + for (int element = 0; element < extent; ++element) { + std::vector args = record_args; + if (field.array_extent) { + args.insert(args.begin(), element); + } + values.push_back(input(args)); + } + } + Func packed("struct_layout_packed"); + packed(records) = pack_struct(record_type_, values); + return {packed}; +} + +std::vector StructLayout::decode(const std::vector &encoded) const { + user_assert(encoded.size() == 1 && encoded[0].outputs() == 1 && + encoded[0].types()[0] == record_type_) + << "StructLayout::decode requires one Func of the exact record type\n"; + const Func &packed = encoded[0]; + std::vector records = component_vars(record_dimensions_, "record"); + std::vector record_args = component_exprs(records); + Expr record = packed(record_args); + std::vector outputs; + outputs.reserve(logical_fields_.size()); + for (const std::string &name : logical_fields_) { + const StructField &physical = physical_field(name); + Func output("struct_layout_" + name); + if (physical.array_extent) { + Var element("element"); + std::vector args = records; + args.insert(args.begin(), element); + output(args) = field(record, name)[element]; + } else { + output(records) = field(record, name); + } + outputs.push_back(output); + } + return outputs; +} + +ApproximationSignature StructLayout::signature() const { + ApproximationSignature sig; + for (const std::string &name : logical_fields_) { + const StructField &f = physical_field(name); + sig.inputs.emplace_back(name, f.type, record_dimensions_ + (f.array_extent ? 1 : 0)); + } + sig.outputs = {{"record", record_type_, record_dimensions_}}; + return sig; +} + +size_t StructLayout::logical_slot(const std::string &name) const { + for (size_t i = 0; i < logical_fields_.size(); ++i) { + if (logical_fields_[i] == name) { + return i; + } + } + user_error << "StructLayout internal field mapping failure\n"; + return 0; +} + +const StructField &StructLayout::physical_field(const std::string &name) const { + for (const StructField &field : record_type_.struct_type()->fields) { + if (field.name == name) { + return field; + } + } + user_error << "StructLayout internal physical field failure\n"; + return record_type_.struct_type()->fields[0]; +} + +std::vector PlanarFieldPack::encode(const std::vector &inputs) const { + user_assert(inputs.size() == 1 && inputs[0].dimensions() == 2) + << "PlanarFieldPack::encode currently requires (element, record)\n"; + const Func &fields = inputs[0]; + Var position("position"), record("record"); + RDom plane(0, planes_, "plane"); + Expr element = plane * positions_ + position; + Expr value = cast(fields(element, record)) & ((1 << field_bits_) - 1); + Func bytes("planar_field_bytes"); + bytes(position, record) = cast(0); + bytes(position, record) = bytes(position, record) | + cast(value << (plane * field_bits_)); + return {bytes}; +} + +std::vector PlanarFieldPack::decode(const std::vector &encoded) const { + user_assert(encoded.size() == 1 && encoded[0].types() == std::vector{UInt(8)} && + encoded[0].dimensions() == 2) + << "PlanarFieldPack::decode currently requires (position, record) bytes\n"; + const Func &bytes = encoded[0]; + Var element("element"), record("record"); + Expr plane = element / positions_; + Expr position = element % positions_; + Func fields("planar_field_values"); + fields(element, record) = cast((bytes(position, record) >> (plane * field_bits_)) & + ((1 << field_bits_) - 1)); + return {fields}; +} + +ApproximationSignature PlanarFieldPack::signature() const { + return {{{"fields", std::nullopt, 2, ApproximationRange(0, (double)((1 << field_bits_) - 1))}}, + {{"bytes", UInt(8), 2}}}; +} + +PlanarFieldPack::PlanarFieldPack(int field_bits, int positions) + : field_bits_(field_bits), positions_(positions), planes_(8 / field_bits) { + user_assert(field_bits_ > 0 && 8 % field_bits_ == 0 && positions_ > 0) + << "Invalid PlanarFieldPack shape\n"; +} + +namespace Internal { + +std::vector little_endian_scalar_encode(Type word_type, const std::vector &inputs) { + user_assert(inputs.size() == 1 && inputs[0].types() == std::vector{word_type}) + << "LittleEndianScalarPack::encode word type mismatch\n"; + const Func &word = inputs[0]; + std::vector records = component_vars(word.dimensions(), "record"); + std::vector record_args(records.begin(), records.end()); + Var byte("byte"); + std::vector args = records; + args.insert(args.begin(), byte); + Func bytes("little_endian_scalar_bytes"); + Expr bits = cast(word_type, word(record_args)); + bytes(args) = cast(bits >> (byte * 8)); + return {bytes}; +} + +std::vector little_endian_scalar_decode(Type word_type, const std::vector &encoded) { + user_assert(encoded.size() == 1 && encoded[0].types() == std::vector{UInt(8)} && + encoded[0].dimensions() >= 2) + << "LittleEndianScalarPack::decode requires byte arrays per record\n"; + const Func &bytes = encoded[0]; + std::vector records = component_vars(bytes.dimensions() - 1, "record"); + std::vector record_args(records.begin(), records.end()); + std::vector pieces; + for (int i = 0; i < word_type.bytes(); ++i) { + std::vector args = record_args; + args.insert(args.begin(), i); + pieces.push_back(bytes(args)); + } + Func word("little_endian_scalar_word"); + word(records) = cast(word_type, concat_bits(pieces)); + return {word}; +} + +ApproximationSignature little_endian_scalar_signature(Type word_type, const ApproximationPorts &inputs) { + if (inputs.size() > 1) { + return ApproximationSignature::unknown(inputs); + } + std::optional dims; + if (inputs.size() == 1) { + dims = inputs[0].dimensions; + } + return {{{inputs.empty() ? "word" : inputs[0].name, word_type, dims}}, + {{"bytes", UInt(8), dims ? std::optional(*dims + 1) : std::nullopt}}}; +} + +} // namespace Internal + +} // namespace Halide diff --git a/src/Approximation.h b/src/Approximation.h new file mode 100644 index 000000000000..2398a5f115d2 --- /dev/null +++ b/src/Approximation.h @@ -0,0 +1,1340 @@ +#ifndef HALIDE_APPROXIMATION_H +#define HALIDE_APPROXIMATION_H + +/** \file + * Defines Approximation, a type-erased handle for lossy, quantified + * Func-to-Func transformations (e.g. a quantize/dequantize round trip), and + * Compose/Parallel/etc., which build larger Approximations out of smaller ones. See + * Func::approximate_by(), which splices such a round trip into an existing + * call graph, and doc/Approximation.md for the design rationale. + */ + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "Func.h" + +namespace Halide { + +struct EncodeResult; +struct DecodeResult; +struct ApproximationResult; +struct ApproximationTraceNode; + +/** A closed interval [lo, hi] of values, with constant double endpoints. This + * is deliberately simpler than Halide::Interval (which holds Exprs and can be + * unbounded on either side): declared value ranges are compile-time constants + * used for documentation and testing, never for code generation. Integers up + * to 2^53 in magnitude are represented exactly. */ +struct ApproximationRange { + double lo = 0, hi = 0; + + ApproximationRange() = default; + ApproximationRange(double lo, double hi) + : lo(lo), hi(hi) { + } + + bool contains(double v) const { + return lo <= v && v <= hi; + } + /** Is every value of `other` in this range? */ + bool contains(const ApproximationRange &other) const { + return lo <= other.lo && other.hi <= hi; + } + bool operator==(const ApproximationRange &other) const { + return lo == other.lo && hi == other.hi; + } +}; + +/** A named slot in an Approximation's interface: one Func that flows into or + * out of a stage. The name identifies the port (for lookups and for + * name-based combinators like Parallel); `type` and `dimensions`, when set, + * are checked against the actual Func at run time. A port with a multi-valued + * (Tuple) Func never has a `type`. + * + * A name identifies a *wire*, not a stage: the port names that reach a stage + * from upstream (or from the caller) are the names it sees, and a stage's + * declared input names are only defaults for when none flow in. The + * declaration is checked (types, dimensions) but never substituted. See + * Approximation::encode(). + * + * `range`, when set, is a declared bound on the port's *values*. Its meaning + * depends on the direction the port is used in: + * + * - on an *input* port of a stage's encode (or a stage's declared inputs): + * a precondition -- the range the stage requires for its declared + * properties (losslessness, error bounds) to hold; + * - on an *output* port of encode: a guarantee -- every encoded value lies + * in the range (e.g. a symmetric int8 quantizer's codes are in [-127, 127]). + * + * Ranges are never checked on the normal encode/decode path, so they cost + * nothing in generated code. They are checked statically (see + * check_ranges() and Approximation::describe()) by comparing each producer's + * guaranteed output range with the consumer's required input range, and + * dynamically by the helpers in tools/halide_approximation_testing.h. */ +struct ApproximationPort { + std::string name; + std::optional type; + std::optional dimensions; + std::optional range; + + ApproximationPort(std::string name, std::optional type = std::nullopt, + std::optional dimensions = std::nullopt, + std::optional range = std::nullopt) + : name(std::move(name)), type(type), dimensions(dimensions), range(range) { + } + ApproximationPort(const char *name, std::optional type = std::nullopt, + std::optional dimensions = std::nullopt, + std::optional range = std::nullopt) + : name(name), type(type), dimensions(dimensions), range(range) { + } +}; + +using ApproximationPorts = std::vector; + +/** The declared interface of an Approximation, in the encode direction: + * `inputs` are consumed by encode() and `outputs` are what it produces. + * Decode is the reverse: it consumes `outputs` and produces `inputs`. + * + * A signature may be *unknown* (`known == false`), meaning the arity of the + * unit could not be determined without running it (e.g. an undeclared + * unit with a multi-Func encode). An unknown signature's `inputs` echo the + * context it was resolved in and its `outputs` are empty. */ +struct ApproximationSignature { + ApproximationPorts inputs, outputs; + bool known = true; + + static ApproximationSignature unknown(ApproximationPorts inputs = {}) { + ApproximationSignature s; + s.inputs = std::move(inputs); + s.known = false; + return s; + } +}; + +/** Approximation is a value-semantic, type-erased handle (in the style of + * std::function) to a lossy, quantified transformation of one or more Funcs' + * values -- e.g. quantize-then-dequantize. Unlike an ordinary schedule + * directive, an Approximation deliberately changes the *value* computed, not + * just how or where it's computed: decode(encode(f)) is expected to + * approximately reproduce f, not exactly reproduce it. + * + * Any type `T` providing const-callable `encode` and `decode` methods + * implicitly converts to an Approximation; there is no base class to derive + * from. Each direction independently may take either ONE of two forms: + * + * \code + * // multi: a vector of Funcs in, a vector of Funcs out + * std::vector encode(const std::vector &) const; + * std::vector decode(const std::vector &) const; + * + * // single: one Func in, one Func out. Handing the unit any other number + * // of inputs is an error. + * Func encode(const Func &) const; + * Func decode(const Func &) const; + * \endcode + * + * The forms may be mixed (e.g. multi encode with single decode). If a type + * offers both a vector and a Func overload for one direction, the vector + * form is used. Types with a non-matching return type (or only one + * direction) do not convert. A minimal unit: + * + * \code + * struct Negate { + * Func encode(const Func &f) const { + * Func g("negated"); + * g(_) = -f(_); + * return g; + * } + * Func decode(const Func &f) const { + * return encode(f); + * } + * }; + * Approximation a = Negate{}; + * \endcode + * + * Each direction may additionally take a *port-aware* form, which also + * receives the ports (names and constraints) of the Funcs that encode() + * consumed -- the original values. In decode() that is the context in which + * the stage's encode ran, not the ports of `encoded`: + * + * \code + * std::vector encode(const std::vector &inputs, + * const ApproximationPorts &input_ports) const; + * std::vector decode(const std::vector &encoded, + * const ApproximationPorts &input_ports) const; + * \endcode + * + * The port-aware form is preferred when present. Combinators use it so that + * they can look ports up by name and hand each child its own context. + * + * A unit may declare its interface with ONE of: + * + * \code + * // static: the same ports whatever the context + * ApproximationSignature signature() const; + * // contextual: given the resolved input ports, return the full signature. + * // `.inputs` normally echoes the given ports, possibly with more constraints. + * // It is called with an empty vector when the inputs are not known. + * ApproximationSignature signature(const ApproximationPorts &inputs) const; + * \endcode + * + * A unit with neither is *undeclared*. Its output ports are named at run + * time: if the output count equals the input count, output i takes input i's + * name (so a single-Func unit preserves its input's name); otherwise the + * outputs are named positionally, "0", "1", .... Declared ports are validated + * against the actual Funcs on every call (see encode()). + * + * Decode never needs to declare anything more: for every stage, the ports of + * decode()'s outputs are named like the ports of its encode()'s inputs, and + * the ports of decode()'s inputs like the ports of encode()'s outputs. + * + * A unit may also list the Approximations it is built from, for describe(): + * + * \code + * std::vector children() const; + * // optional: the input ports each child would receive when this unit's + * // encode receives `inputs`, parallel to children(); it is what lets + * // describe() and check_ranges() resolve the children's signatures + * std::vector child_inputs(const ApproximationPorts &inputs) const; + * \endcode + * + * A unit may declare how accurate its round trip is, with ONE of: + * + * \code + * // A per-element bound on |decode(encode(x)) - x|, as a Func with the same + * // arguments as inputs[0] (any numeric type). `encoded` are the Funcs that + * // encode() produced from `inputs`. + * Func error_bound(const std::vector &inputs, const std::vector &encoded) const; + * // The round trip is exact (the bound is zero). + * bool lossless() const; + * \endcode + * + * Both are claims that hold *when every input port's declared range (its + * precondition) is respected*; see Approximation::error_bound() and + * lossless(). They are never enforced at run time. + * + * See also Pointwise for elementwise units. The handle stores a decayed copy + * of the unit, so methods are always invoked on a const object. (Units + * needing mutable state must hold it in `mutable` members or behind a + * pointer.) + * + * The handle's own encode()/decode() always take and return a *vector* of + * Funcs, even though the common case (a leaf Approximation like a plain + * quantizer) only ever uses one. This is what makes Compose and Parallel below + * possible: a composed Approximation's inner stage can produce multiple Funcs + * (e.g. a quantized-values Func plus a separate scale Func), and the next + * stage needs to be able to consume all of them, or select just one to act + * on. + * + * A unit only returns its output Funcs. The framework discovers the rest: + * whatever intermediate Funcs the unit defined along the way (e.g. a + * per-block reduction) are found by walking the outputs' definitions, and + * are reported in EncodeResult::intermediates / DecodeResult::intermediates + * (see Approximation::encode()). A unit that calls other Approximation + * handles inside its own encode()/decode() (as Compose does) gets those + * calls traced automatically. + * + * Identity semantics: a copy of a handle is the *same* stage (same_as() is + * true), while converting a plain unit to an Approximation twice produces two + * *distinct* stages, even if the units compare equal. Every stage invoked + * through a handle -- including those invoked from inside another unit's + * encode()/decode() -- is recorded automatically in the result's + * `stage_outputs` (children first, then the stage itself), so to find a stage's outputs + * later, hold onto a handle to it and hand copies of that handle to the + * combinators: + * + * \code + * Approximation qh = LittleEndianScalarPack{}; + * Compose scheme{BlockReshape{32}, qh}; + * ApproximationResult r = f.approximate_by(scheme, {g}); + * Func bytes = r.decoded_by(qh); + * \endcode + * + * An Approximation makes no claim about *where* or *when* encode/decode are + * computed relative to the rest of a pipeline (offline vs fused inline, + * compute_root vs compute_at) -- that is a scheduling decision, orthogonal + * to the semantics defined here. Concretely: the same Approximation can be + * used with encode() computed once, offline, ahead of any other stage (a + * static weight quantizer) or fused into a producer's inner loop and + * recomputed on every call (dynamic activation requantization) -- nothing + * about the interface favors one over the other. See Func::approximate_by() + * for splicing an Approximation into an existing call graph. */ +class Approximation { + enum class SignatureForm { None, + Static, + Contextual }; + + struct Concept { + virtual ~Concept() = default; + virtual std::vector encode(const std::vector &inputs, + const ApproximationPorts &input_ports) const = 0; + virtual std::vector decode(const std::vector &encoded, + const ApproximationPorts &input_ports) const = 0; + virtual std::string default_label() const = 0; + virtual SignatureForm signature_form() const = 0; + virtual ApproximationSignature declared_signature(const ApproximationPorts &inputs) const = 0; + virtual bool encode_is_single() const = 0; + virtual std::vector children() const = 0; + virtual std::vector child_inputs(const ApproximationPorts &inputs) const = 0; + virtual Func error_bound(const std::vector &inputs, const std::vector &encoded) const = 0; + virtual bool lossless() const = 0; + }; + + // The shared identity of a stage: every copy of a handle points at one + // State, so the label lives here and is seen by all copies. + struct State { + std::unique_ptr impl; + std::string label; + }; + + /** A readable name for a unit type: demangled, with "Halide::" and + * anonymous-namespace qualifiers stripped. */ + static std::string type_label(const char *pretty_function); + + // Captures the compiler's pretty function name, which spells out T; this + // avoids depending on RTTI. + template + static const char *pretty_type_name() { +#if defined(_MSC_VER) && !defined(__clang__) + return __FUNCSIG__; +#else + return __PRETTY_FUNCTION__; +#endif + } + + template + struct Model; + + enum class Form { None, + Ported, + Multi, + Single }; + + template + using encode_call_t = decltype(std::declval().encode(std::declval())); + template + using decode_call_t = decltype(std::declval().decode(std::declval())); + + template + struct form_of : std::integral_constant {}; + + // The vector-argument call is tried first, so a type with both vector + // and Func overloads uses the vector form. + template + struct form_of>>> + : std::integral_constant {}; + + template + struct form_of> && + std::is_same_v>> + : std::integral_constant {}; + + template + struct detect_encode_vec { + using type = void; + }; + template + struct detect_encode_vec &>>> { + using type = std::decay_t &>>; + }; + template + struct detect_encode_one { + using type = void; + }; + template + struct detect_encode_one>> { + using type = std::decay_t>; + }; + template + struct detect_decode_vec { + using type = void; + }; + template + struct detect_decode_vec &>>> { + using type = std::decay_t &>>; + }; + template + struct detect_decode_one { + using type = void; + }; + template + struct detect_decode_one>> { + using type = std::decay_t>; + }; + + template + struct has_name : std::false_type {}; + template + struct has_name().name()))>> + : std::true_type {}; + + template + struct is_ported_encode : std::false_type {}; + template + struct is_ported_encode().encode( + std::declval &>(), + std::declval()))>, + std::vector>>> : std::true_type {}; + template + struct is_ported_decode : std::false_type {}; + template + struct is_ported_decode().decode( + std::declval &>(), + std::declval()))>, + std::vector>>> : std::true_type {}; + + template + struct has_static_signature : std::false_type {}; + template + struct has_static_signature().signature())>, + ApproximationSignature>>> : std::true_type {}; + template + struct has_contextual_signature : std::false_type {}; + template + struct has_contextual_signature().signature( + std::declval()))>, + ApproximationSignature>>> : std::true_type {}; + + template + struct has_children : std::false_type {}; + template + struct has_children().children())>, + std::vector>>> : std::true_type {}; + template + struct has_child_inputs : std::false_type {}; + template + struct has_child_inputs().child_inputs( + std::declval()))>, + std::vector>>> : std::true_type {}; + + template + struct has_error_bound : std::false_type {}; + template + struct has_error_bound().error_bound( + std::declval &>(), + std::declval &>()))>, + Func>>> : std::true_type {}; + template + struct has_lossless : std::false_type {}; + template + struct has_lossless().lossless()), bool>>> : std::true_type {}; + + template + static constexpr Form encode_form = + is_ported_encode::value ? Form::Ported : + form_of::type, + typename detect_encode_one::type>::value; + template + static constexpr Form decode_form = + is_ported_decode::value ? Form::Ported : + form_of::type, + typename detect_decode_one::type>::value; + + template + using enable_if_unit = std::enable_if_t< + !std::is_base_of_v> && + encode_form> != Form::None && decode_form> != Form::None>; + + static void check_single_input(const std::vector &inputs, const char *direction); + +public: + /** Construct an undefined handle. */ + Approximation() = default; + + /** Wrap a copy of `unit` as a new stage. */ + template> + Approximation(T &&unit); + + /** Wrap a copy of `unit` as a new stage with an explicit label. */ + template> + Approximation(T &&unit, std::string label); + + bool defined() const { + return state_ != nullptr; + } + + /** Do these handles refer to the same stage? True for copies of one + * handle; false for two separate conversions of equal units. */ + bool same_as(const Approximation &other) const { + return state_ == other.state_; + } + + /** Set this stage's label and return this handle. The label lives in the + * state shared by every copy of the handle (they are all the same stage), + * so this affects all existing copies too -- and the result is same_as() + * this handle. Typical use: `Approximation q = Approximation(unit).labelled("q");` + * or `Approximation q(unit, "q");`. */ + Approximation labelled(std::string label) const; + + /** A human-readable name for this stage, used when printing traces. It is + * the label set by labelled() (or the constructor), if any; otherwise the + * unit's `std::string name() const` if it has one; otherwise the unit's + * type name with "Halide::" qualifiers stripped (e.g. "Compose", + * "LittleEndianScalarPack"). Empty for an undefined handle. */ + std::string label() const; + + /** Produce the encoded form of `inputs`. EncodeResult::encoded's + * elements are not required to have the same type, dimensionality, or + * count as `inputs` -- an Approximation is free to choose a packed + * representation (a single opaque byte buffer, fields recovered via + * reinterpret<>() inside decode) or a planar one (multiple typed Funcs, + * one per field). Either is legitimate; the framework does not + * decide. + * + * After the unit returns, the framework computes this stage's + * `intermediates`: every Func reachable from the outputs' definitions + * (pure, update, and extern-argument references) without passing through + * one of `inputs`, excluding the inputs and outputs themselves, in + * topological order (producers before consumers; ties broken by name). + * This includes pure-only Funcs; callers filter as they see fit. Funcs + * the unit references from outside (e.g. a shared lookup table Func + * defined elsewhere) are reachable and not inputs, so they are reported + * too. + * + * The outermost handle call on a thread opens a trace, and every handle + * call nested inside it (directly or from inside any unit's + * encode()/decode()) appends a record to it, children first, then the + * stage itself. `stage_outputs` holds exactly the records appended + * during this call. Encode and decode share one trace: a decode() call + * made from inside an encode() (unusual) is recorded in that encode's + * `stage_outputs`. + * + * Ports: `input_ports` name the `inputs`, and a name that flows in like + * this is the one the stage uses, whatever its declared signature says. + * If `input_ports` is non-empty, its size must equal `inputs.size()`. If + * empty, the unit's declared signature supplies default names (when its + * input count matches), else they are positional: "0", "1", .... Every + * input Func is then checked against its port's `type` and `dimensions` + * (when set) -- both the given ports and the unit's declared inputs -- and + * a mismatch is a user_error naming this stage, the direction, the port, + * and expected vs actual. The output ports are the declared signature's + * outputs, resolved for these inputs (which must match the output count, + * and are validated like the inputs) or, for an undeclared or + * unknown-signature unit, follow the naming rule described on + * Approximation. They are returned in EncodeResult::encoded_ports, with + * unset types and dimensions filled in from the actual Funcs, so that + * they can be handed to the next stage. */ + EncodeResult encode(const std::vector &inputs, + const ApproximationPorts &input_ports = {}) const; + + /** Reconstruct an approximation of the original Func(s) from their + * encoded form. See DecodeResult for the constraint on `decoded`'s + * size, which depends on how this Approximation is used. + * + * `input_ports` are the ports encode() was given (or empty, if it was + * given none, as when a decode runs on its own after sever + * severs the encode). They are the *context*: the encoded ports and the + * decoded ports are both derived from them statically, through the + * declared signature (as describe() does), so that a stand-alone decode + * agrees with the encode that produced its inputs. Concretely, the + * `encoded` Funcs are checked against, and named after, the outputs of + * signature(input_ports), and the decoded Funcs are named after + * `input_ports` (the defaults, if empty): decode's output port i is named + * like encode's input port i. For an undeclared unit whose signature is + * unknown, the encoded ports are named by the naming rule and the + * decoded ports follow `input_ports` if their count matches. */ + DecodeResult decode(const std::vector &encoded, + const ApproximationPorts &input_ports = {}) const; + + /** The signature of this stage in the encode direction, resolved for + * inputs named `inputs` (empty if unknown). Nothing is run. + * + * - static signature: returned as is, except that when `inputs` is + * non-empty and has the declared size, the input ports take the names + * of `inputs` (the declared names are defaults, never substitutes); + * - contextual signature: computed from `inputs`, with the same rule + * for the input names; + * - undeclared single-Func unit: one input (`inputs`, or "0" if none + * were given) and one output with the same name; + * - undeclared multi-Func unit: unknown (see ApproximationSignature). */ + ApproximationSignature signature(const ApproximationPorts &inputs = {}) const; + + /** A per-element upper bound on |decode(encode(x)) - x|, given the + * `inputs` handed to encode() and the `encoded` Funcs it returned, valid + * when the inputs respect the declared input ranges. The result has the + * same arguments as inputs[0] and a numeric type. It comes from the + * unit's `error_bound()` if it has one, else a zero Func if the unit is + * lossless(), else an undefined Func (no bound is declared). The bound + * is a claim, not enforced; see tools/halide_approximation_testing.h to check it. */ + Func error_bound(const std::vector &inputs, const std::vector &encoded) const; + + /** Does the unit declare its round trip exact (under its input + * preconditions)? A unit that only has an `error_bound()` is not + * lossless() even if that bound happens to be zero. */ + bool lossless() const; + + /** Render this stage's structure without running anything: one line per + * stage, `label (inputs) -> (outputs)`, where a port prints as + * `name: type xN in [lo, hi]` (`N` being its dimensionality; unset parts + * are left out; the range is a precondition on inputs and a guarantee on + * outputs), followed by the stage's children (see the `children()` unit + * method), indented by two spaces. `inputs` is the context in which the + * signature is resolved. When a stage's context is known, each input + * whose declared range is not guaranteed by the context is flagged on the + * following line, e.g. `! input 'value': [0, 15] not within [0, 7]` + * (or `... not guaranteed` when the context declares no range). */ + std::string describe(const ApproximationPorts &inputs = {}) const; + +private: + ApproximationPorts resolve_inputs(const std::vector &inputs, const ApproximationPorts &given) const; + ApproximationPorts resolve_encoded(const std::vector &encoded, const ApproximationSignature &sig, + const ApproximationPorts &input_ports) const; + ApproximationPorts output_ports(const std::vector &outputs, const ApproximationSignature &sig, + const ApproximationPorts &input_ports, bool encode_direction) const; + void describe_to(std::string &out, const ApproximationPorts &inputs, int depth) const; + void range_issues_to(std::vector &out, const ApproximationPorts &inputs, + const std::string &path) const; + friend std::vector check_ranges(const Approximation &, const ApproximationPorts &); + + std::shared_ptr state_; +}; + +/** One handle call (encode or decode) in an execution trace. `children` are + * the handle calls that started and finished during this one, in the order + * they were invoked; a node completes after all of its children. */ +struct ApproximationTraceNode { + Approximation stage; + /** stage.label() when the call finished. */ + std::string label; + /** The stage's output Funcs for this call. */ + std::vector ports; + /** The names of `ports` (parallel to it). */ + std::vector port_names; + /** The Funcs discovered for this stage alone (see Approximation::encode). */ + std::vector intermediates; + std::vector children; + /** The Funcs this call consumed (the stage's inputs) and their port + * names (parallel to it). */ + std::vector inputs; + std::vector input_names; +}; + +/** The ports produced by one stage during encode or decode, plus the + * intermediate Funcs discovered for that stage alone. This trace is + * supplemental scheduling metadata; it does not alter the signature + * contract. */ +struct ApproximationStageOutputs { + Approximation stage; + std::vector ports; + /** The names of `ports` (parallel to it). */ + std::vector port_names; + std::vector intermediates; +}; + +/** The result of Approximation::encode(): the Func(s) that make up the + * signature contract other code is expected to consume, plus the + * intermediate Funcs discovered between the inputs and those outputs (e.g. + * per-block reduction Funcs), which have no meaning outside scheduling but + * must still be scheduled by whoever calls encode(). */ +struct EncodeResult { + std::vector encoded; + /** The ports of `encoded` (parallel to it), ready to hand to the next + * stage's encode() or to decode(). */ + ApproximationPorts encoded_ports; + std::vector intermediates; + /** The trace of this call as a flat list: the post-order flattening of + * `trace` (children first, then the stage itself). */ + std::vector stage_outputs; + /** The trace of this call as a tree; its root is this call's stage. */ + ApproximationTraceNode trace; +}; + +/** The result of Approximation::decode(): decoded is the round-trip + * replacement for whatever Func(s) were originally encoded, plus the + * discovered scheduling-only intermediates. When an Approximation is used + * directly with Func::approximate_by(), decoded must contain exactly one + * Func; when it's used as one stage of a larger Compose/Parallel chain, + * decoded may contain however many Funcs the next stage down expects. */ +struct DecodeResult { + std::vector decoded; + /** The ports of `decoded` (parallel to it). */ + ApproximationPorts decoded_ports; + std::vector intermediates; + std::vector stage_outputs; + ApproximationTraceNode trace; +}; + +template +struct Approximation::Model final : Approximation::Concept { + T unit; + + template, Model>>> + explicit Model(U &&u) + : unit(std::forward(u)) { + } + + std::vector encode(const std::vector &inputs, + const ApproximationPorts &input_ports) const override { + if constexpr (encode_form == Form::Ported) { + return unit.encode(inputs, input_ports); + } else if constexpr (encode_form == Form::Multi) { + return unit.encode(inputs); + } else { + check_single_input(inputs, "encode"); + return {unit.encode(inputs[0])}; + } + } + + std::vector decode(const std::vector &encoded, + const ApproximationPorts &input_ports) const override { + if constexpr (decode_form == Form::Ported) { + return unit.decode(encoded, input_ports); + } else if constexpr (decode_form == Form::Multi) { + return unit.decode(encoded); + } else { + check_single_input(encoded, "decode"); + return {unit.decode(encoded[0])}; + } + } + + std::string default_label() const override { + if constexpr (has_name::value) { + return std::string(unit.name()); + } else { + return type_label(pretty_type_name()); + } + } + + SignatureForm signature_form() const override { + if constexpr (has_contextual_signature::value) { + return SignatureForm::Contextual; + } else if constexpr (has_static_signature::value) { + return SignatureForm::Static; + } else { + return SignatureForm::None; + } + } + + ApproximationSignature declared_signature(const ApproximationPorts &inputs) const override { + if constexpr (has_contextual_signature::value) { + return unit.signature(inputs); + } else if constexpr (has_static_signature::value) { + return unit.signature(); + } else { + return ApproximationSignature::unknown(inputs); + } + } + + bool encode_is_single() const override { + return encode_form == Form::Single; + } + + std::vector children() const override { + if constexpr (has_children::value) { + return unit.children(); + } else { + return {}; + } + } + + std::vector child_inputs(const ApproximationPorts &inputs) const override { + if constexpr (has_child_inputs::value) { + return unit.child_inputs(inputs); + } else { + return std::vector(children().size()); + } + } + + Func error_bound(const std::vector &inputs, const std::vector &encoded) const override { + if constexpr (has_error_bound::value) { + return unit.error_bound(inputs, encoded); + } else { + return Func(); + } + } + + bool lossless() const override { + if constexpr (has_lossless::value) { + return unit.lossless(); + } else { + return false; + } + } +}; + +template +Approximation::Approximation(T &&unit) + : Approximation(std::forward(unit), std::string()) { +} + +template +Approximation::Approximation(T &&unit, std::string label) + : state_(std::make_shared()) { + state_->impl = std::make_unique>>(std::forward(unit)); + state_->label = std::move(label); +} + +/** The result of Func::approximate_by(): the primary replacement Func + * (already spliced into every Func in `consumers`), plus every + * intermediate Func discovered by encode()/decode() along the way that needs + * scheduling (compute_root, compute_at, etc.) -- none of `intermediates` are + * part of the Approximation's signature contract, but Halide still + * requires Funcs with update definitions to be scheduled, and the fusion + * patterns described on Approximation above (e.g. compute_at-ing the + * encoded Func into a producer) are only possible if the caller has a + * Func to schedule. `intermediates` is `encoded`, then the encode side's + * discovered intermediates, then the decode side's, without duplicates, and + * never contains the original Func or `replacement`. Like all discovered + * intermediates, it may include Funcs the units referenced from outside. */ +struct ApproximationResult { + Func replacement; + /** The Func(s) produced by encode() -- the signature-contract boundary + * between the original values and their approximated form (e.g. a + * quantizer's packed byte buffer). This is a subset of `intermediates` (kept + * there too, so existing code that schedules everything in `intermediates` + * doesn't need to change), broken out separately so callers can act on + * exactly this boundary -- e.g. Pipeline::sever(result.encoded) + * -- without calling Approximation::encode() themselves. */ + std::vector encoded; + /** The ports of `encoded` (parallel to it). */ + ApproximationPorts encoded_ports; + std::vector intermediates; + std::vector encoded_stage_outputs; + std::vector decoded_stage_outputs; + /** The encode and decode traces as trees; the roots are the stages + * invoked by approximate_by() itself. */ + ApproximationTraceNode encode_trace; + ApproximationTraceNode decode_trace; + + /** Return the given output port of `stage` (found by same_as()), or an + * undefined Func if the port is out of range. It is an error if `stage` + * was not invoked in this direction, or was invoked more than once (the + * lookup would be ambiguous). */ + Func encoded_by(const Approximation &stage, size_t port = 0) const; + Func decoded_by(const Approximation &stage, size_t port = 0) const; + + /** As above, but find the port by name. It is an error if `stage` has no + * port of that name in this direction (the message lists the ports it + * has), or if the name is ambiguous. */ + Func encoded_by(const Approximation &stage, const std::string &port) const; + Func decoded_by(const Approximation &stage, const std::string &port) const; + + /** Every Func that is an output port of some stage in either direction, + * deduplicated by name, in trace order (encode side first, each side + * post-order: children before parents). Callers can use it to schedule + * stage boundaries (e.g. compute_root them alongside reductions). + * `replacement` -- the decode root's output, already spliced into the + * consumers -- is excluded. A stage that passes an input through + * (Identity, or Parallel around one) reports it as a port, so the original Func may appear + * if such a stage sits at the very inside of the encode chain. */ + std::vector stage_ports() const; + + /** Is `f` (matched by name) one of stage_ports()? */ + bool is_stage_port(const Func &f) const; +}; + +/** Statically check declared value ranges through `a`, without running + * anything. For every stage whose input context is known (the root's is + * `inputs`, which is skipped when empty; a child's comes from + * `child_inputs()`), each input port with a declared range (a precondition) + * is compared with the range the context guarantees: + * + * - if both are declared and the guarantee is not within the precondition, + * the diagnostic reads `path: input 'p': [0, 15] not within [0, 7]`; + * - if the guarantee is undeclared, `path: input 'p': requires [0, 7], but + * the producer declares no range`. + * + * `path` is the chain of stage labels from the root. An empty result means + * every declared precondition is statically guaranteed. Ranges only flow + * through units that declare them, so unknown is common; unknown is reported, + * never assumed satisfied. */ +std::vector check_ranges(const Approximation &a, const ApproximationPorts &inputs = {}); + +/** Print a trace as an indented tree, one line per call: the label, then + * `-> ` and the comma-separated ports as `port name=Func name`. A non-empty intermediates + * list follows on its own line as `intermediates: a, b`, indented under its + * stage, then the children in invocation order. */ +std::ostream &operator<<(std::ostream &stream, const ApproximationTraceNode &node); + +/** Print Approximation::describe(). */ +std::ostream &operator<<(std::ostream &stream, const Approximation &approximation); + +/** Print both directions of an ApproximationResult under `encode:` and + * `decode:` headers. */ +std::ostream &operator<<(std::ostream &stream, const ApproximationResult &result); + +/** Sequentially composes any number of Approximations into a pipeline. The + * stages are listed in *encode order*, innermost first: encode() runs them + * first to last, on the original inputs, feeding each stage's encoded output + * to the next; decode() runs the mirror image, last to first. So `stages[0]` + * is closest to the original values, and `stages.back()` is the one whose + * encode() output is this Compose's own encoded result, and whose decode() + * input is this Compose's own encoded argument. + * + * Each stage is held as an Approximation handle (plain units convert + * implicitly), so pass a named handle for any stage you want to look up + * later with ApproximationResult::encoded_by()/decoded_by(): + * + * \code + * Compose scheme{ + * BlockReshape{block_size}, + * SymmetricAffineQuantize{block_size, qmax, rounding, anchor}, + * Parallel{{"codes", Fp8Pack{}}, {"scale", Fp16Pack{}}}, + * StructPack{...}, + * }; + * \endcode + * + * Each stage sees the port names that the previous stage's encode produced, + * and (in decode) the ports are derived statically from the context in the + * same way, so a stage's decode output ports are named like its encode + * input ports. */ +struct Compose { + explicit Compose(std::vector stages) + : stages(std::move(stages)) { + } + + template, + std::is_convertible, + std::is_convertible...>>> + Compose(A &&a, B &&b, Rest &&...rest) + : stages{Approximation(std::forward(a)), Approximation(std::forward(b)), + Approximation(std::forward(rest))...} { + } + + /** Each stage receives the ports the previous stage produced. */ + std::vector encode(const std::vector &inputs, const ApproximationPorts &input_ports) const; + + /** Each stage receives the context that its encode had, found by + * threading the stages' signatures from `input_ports`. */ + std::vector decode(const std::vector &encoded, const ApproximationPorts &input_ports) const; + + /** Chains the stages' signatures from the first. Unknown if any stage's + * signature is unknown. */ + ApproximationSignature signature(const ApproximationPorts &inputs) const; + + /** The stages in encode order, each with the ports it would receive. */ + std::vector children() const; + std::vector child_inputs(const ApproximationPorts &inputs) const; + + /** A composition of lossless stages is lossless (under their + * preconditions, see check_ranges()). No error bound is declared for + * lossy compositions: bounds do not compose without knowing how errors + * propagate. */ + bool lossless() const; + + std::vector stages; +}; + +/** A product combinator: applies different Approximations to different parts + * of a vector of Funcs, side by side. It comes in two forms. + * + * *Positional*: `Parallel{a, b, c}`. Child i is applied to a consecutive + * slice of the Funcs, whose width is the number of inputs in the child's + * signature in encode() (and of outputs in decode(), i.e. the number of + * encoded ports), or one if the child's signature is unknown. The slices + * must exactly cover the Funcs, or it is an error. Use Identity to pass one + * Func through. + * + * *Named*: `Parallel{{"codes", a}, {"scale", b}}`. Each entry routes the port + * called `"codes"` to its child (an error if there is no such port, or more + * than one, or if two entries name the same port). Ports that are not + * mentioned pass through unchanged, in place. A child's outputs replace its + * port in place, so a child may expand one port into several in encode(), and + * in decode() collapses them back into one. Naming the ports requires the + * context to have names: a by-name Parallel has no known signature without + * input ports, which is the case for the first stage of a Compose only if + * given ports explicitly. + * + * The two forms cannot be mixed. In both, the Funcs the children produce, + * in order, are Parallel's outputs, and the children see (and produce) port + * names as they are, so names flow through untouched. + * + * Parallel is lossless if all of its children are. It declares no error bound, + * since it would have to know how the children's errors combine. */ +struct Parallel { + /** One by-name entry. */ + struct Entry { + Entry(std::string port, Approximation child) + : port(std::move(port)), child(std::move(child)) { + } + std::string port; + Approximation child; + }; + + /** Positional. */ + explicit Parallel(std::vector children) + : children_(std::move(children)) { + } + + template, + std::is_convertible, + std::is_convertible...>>> + Parallel(A &&a, B &&b, Rest &&...rest) + : children_{Approximation(std::forward(a)), Approximation(std::forward(b)), + Approximation(std::forward(rest))...} { + } + + /** Named. */ + Parallel(std::initializer_list entries); + + std::vector encode(const std::vector &inputs, const ApproximationPorts &input_ports) const; + std::vector decode(const std::vector &encoded, const ApproximationPorts &input_ports) const; + + /** Unknown if the children cannot be located in `inputs` or any of their + * signatures is unknown. Without `inputs`, a positional Parallel uses + * each child's own defaults. */ + ApproximationSignature signature(const ApproximationPorts &inputs) const; + + /** The children, with the ports each would receive. */ + std::vector children() const; + std::vector child_inputs(const ApproximationPorts &inputs) const; + + /** Lossless if every child is. */ + bool lossless() const; + +private: + // How the encode-side ports are divided among the children, in port + // order: `child` (an index into children_, or -1 for a port that passes + // through) takes [begin, begin + width) of the ports, with `context` as + // their ports (empty if unknown), and produces `outputs` encoded ports. + struct Segment { + int child; + size_t begin, width; + ApproximationPorts context; + size_t outputs; + }; + bool plan(const ApproximationPorts &inputs, std::vector &segments, std::string *problem) const; + + std::vector children_; + std::vector ports_; +}; + +/** Routes encode() to one Approximation and decode() to another, taking each + * direction from a *different* source. This is the deliberate backdoor out of + * the structural guarantee Compose provides. + * + * Every Approximation is meant to be an approximate identity, factored into a + * decode-after-encode pair (decode(encode(f)) ~= f). Compose preserves that by + * construction: it interleaves its stages' encode()s and decode()s in mirror + * order, so the composed round trip (d1 . d2) . (e2 . e1) is *guaranteed* to be + * an approximate identity for the same structural reason each stage is -- the + * two halves provably come from one stage list. TrustedInverse pairs an encode + * and a decode from unrelated Approximations, so nothing structural guarantees + * they compose to an identity: the caller is *trusted* to have supplied a true + * inverse pair. Hence the name -- "trusted" as in "taken on trust", not "known + * safe". + * + * The motivating case: a scheme whose forward map (quantize) is an opaque + * offline black box -- a per-block codeword search, a transcendental scale fit, + * typically an extern call -- that no composition of Halide Funcs reproduces + * bit-for-bit, but whose reverse map (dequantize) *is* an ordinary Compose of + * invertible primitives. Compose can't express that pairing; TrustedInverse + * can, keeping the decode side a clean composition while the encode side is + * whatever opaque Approximation actually produces the encoded form: + * + * \code + * TrustedInverse{ + * ExternQuantize{"q4_k_quantize_via_ggml"}, // encode(): values -> bytes + * Compose{ // decode(): bytes -> values + * BlockReshape{block_size}, ..., Parallel{...}, StructPack{...}, + * }, + * }; + * \endcode + * + * The unused half of each side is never called (here, the ExternQuantize's + * decode() and the Compose's encode()); supplying an Approximation whose + * relevant half is a stub is expected. */ +struct TrustedInverse { + TrustedInverse(Approximation encoder, Approximation decoder) + : encoder(std::move(encoder)), decoder(std::move(decoder)) { + } + + std::vector encode(const std::vector &inputs, const ApproximationPorts &input_ports) const; + std::vector decode(const std::vector &encoded, const ApproximationPorts &input_ports) const; + + /** The encoder's signature, in both directions: the encoder defines the + * encoded representation, and the decoder is trusted to consume it. */ + ApproximationSignature signature(const ApproximationPorts &inputs) const; + + /** The encoder, then the decoder (which is described without context, + * since it runs in the other direction). */ + std::vector children() const; + std::vector child_inputs(const ApproximationPorts &inputs) const; + + Approximation encoder, decoder; +}; + +/** Picks one of two Approximations at construction time based on `cond`, + * keeping only the chosen handle (so it can be looked up by that handle). */ +struct Choose { + Choose(bool cond, Approximation if_true, Approximation if_false) + : chosen(cond ? std::move(if_true) : std::move(if_false)) { + } + + std::vector encode(const std::vector &inputs, const ApproximationPorts &input_ports) const; + std::vector decode(const std::vector &encoded, const ApproximationPorts &input_ports) const; + + /** The chosen stage's signature. */ + ApproximationSignature signature(const ApproximationPorts &inputs) const; + + /** Just the chosen stage. */ + std::vector children() const; + std::vector child_inputs(const ApproximationPorts &inputs) const; + + /** Lossless if the chosen stage is. */ + bool lossless() const; + + Approximation chosen; +}; + +/** An elementwise unit: encode() and decode() each map every value of a + * single input Func through a user-supplied function, producing a Func with + * the same dimensionality (`out(vs) = fn(in(vs))`). + * + * The functions take an Expr and return an Expr (the common case, for + * single-valued Funcs), or take and return a `std::vector` (one element + * per Tuple output of the input Func, for Tuple-valued Funcs). Pass + * lambdas with concrete parameter types, not `auto`. The input is not type + * checked unless with_types() declares its type. + * + * The output Funcs are named `name + "_encode"` and `name + "_decode"`; the + * second constructor lets the caller pick both names exactly (and the + * prefix of the pure Var names, which are numbered by dimension), so that a + * unit built on Pointwise can keep its Func names stable. + * + * \code + * Approximation a = Pointwise{"scale", + * [](Expr x) { return x * 2; }, + * [](Expr x) { return x / 2; }}; + * \endcode + * + * A cast is a one-liner: + * + * \code + * Approximation to_f16 = Pointwise{"f16", + * [](Expr x) { return cast(x); }, + * [](Expr x) { return cast(x); }} + * .with_types(Float(32), Float(16)); + * \endcode + * + * By default a Pointwise declares only its arity: one Func in, one out, with + * the name and dimensionality of the input port. The with_*() methods + * declare more of the signature (see Approximation::signature()), and + * claim losslessness. + * + * Besides converting to Approximation, Pointwise's encode(Func) and + * decode(Func) may be called directly. */ +struct Pointwise { + using ExprFn = std::function; + using TupleFn = std::function(const std::vector &)>; + + Pointwise(const std::string &name, ExprFn encode_fn, ExprFn decode_fn); + Pointwise(const std::string &name, TupleFn encode_fn, TupleFn decode_fn); + Pointwise(std::string encode_name, std::string decode_name, ExprFn encode_fn, ExprFn decode_fn, + std::string var_prefix = "pw"); + Pointwise(std::string encode_name, std::string decode_name, TupleFn encode_fn, TupleFn decode_fn, + std::string var_prefix = "pw"); + + Func encode(const Func &input) const; + Func decode(const Func &encoded) const; + + /** A copy that also declares an error bound: `bound(x)` is an upper + * bound on |decode(encode(x)) - x| as a function of the *original* value + * x (single-valued Funcs only). Without this, error_bound() returns an + * undefined Func. */ + Pointwise with_error_bound(ExprFn bound) const; + Func error_bound(const std::vector &inputs, const std::vector &encoded) const; + + /** A copy that declares the type of its input (what encode consumes) and + * of its output (what encode produces); both are checked against the + * actual Funcs. */ + Pointwise with_types(Type input_type, Type output_type) const; + + /** A copy that declares value ranges: `input_range` is the precondition + * on the input for the unit's declared properties to hold (e.g. + * losslessness), and `output_range` is a guarantee on the encoded + * values (see ApproximationPort::range). */ + Pointwise with_ranges(ApproximationRange input_range, ApproximationRange output_range) const; + + /** A copy that claims to be lossless (for inputs within the declared + * input range). */ + Pointwise with_lossless(bool is_lossless = true) const; + + /** The declared signature; see the with_*() methods. Undeclared parts are + * left unknown. */ + ApproximationSignature signature(const ApproximationPorts &inputs) const; + + bool lossless() const { + return lossless_; + } + +private: + std::string encode_name, decode_name, var_prefix; + TupleFn encode_fn, decode_fn; + ExprFn bound_fn; + std::optional input_type, output_type; + std::optional input_range, output_range; + bool lossless_ = false; +}; + +/** Passes Funcs (and their port names) through unchanged in both directions. */ +struct Identity { + std::vector encode(const std::vector &inputs) const; + std::vector decode(const std::vector &encoded) const; + + /** Echoes `inputs`; unknown if there are none. */ + ApproximationSignature signature(const ApproximationPorts &inputs) const; + + bool lossless() const { + return true; + } +}; + +/** Reorders Funcs: encode() outputs `inputs[permutation[i]]` at position i, + * and decode() inverts that. The output port names are permuted the same + * way. Without a context (no input ports), the inputs are named + * positionally. */ +struct Permute { + explicit Permute(std::vector permutation) + : forward(std::move(permutation)) { + backward.resize(forward.size()); + for (int i = 0; i < (int)forward.size(); i++) { + backward[forward[i]] = i; + } + } + + std::vector encode(const std::vector &inputs) const; + std::vector decode(const std::vector &encoded) const; + + ApproximationSignature signature(const ApproximationPorts &inputs) const; + + bool lossless() const { + return true; + } + + std::vector forward, backward; +}; + +/** Losslessly reshape a flat row into fixed-size records. In block-indexed + * mode the flat side is `(within, record)` rather than a single flat index. */ +struct BlockReshape { + explicit BlockReshape(int block_size, bool block_indexed = false) + : extents_{block_size}, block_indexed_(block_indexed) { + } + explicit BlockReshape(std::vector extents, bool block_indexed = false) + : extents_(std::move(extents)), block_indexed_(block_indexed) { + } + + std::vector encode(const std::vector &inputs) const; + std::vector decode(const std::vector &encoded) const; + + /** values (flat) <-> blocks (one extra leading dimension per extent). */ + ApproximationSignature signature() const; + + /** Pure re-indexing: exact for any values (given the flat extent is a + * multiple of the block size). */ + bool lossless() const { + return true; + } + +private: + std::vector extents_; + bool block_indexed_; + + int block_size() const; + std::vector block_vars() const; +}; + +/** Map consecutive logical Func slots to named fields of an exact struct + * type. Scalar fields have `record_dimensions` dimensions; array fields have + * an additional leading element dimension. */ +struct StructLayout { + StructLayout(Type record_type, std::vector logical_fields, + int record_dimensions = 1); + + std::vector encode(const std::vector &inputs) const; + std::vector decode(const std::vector &encoded) const; + + /** One input per logical field, named after it, and one `record` output. */ + ApproximationSignature signature() const; + + /** Fields are stored bit-for-bit. */ + bool lossless() const { + return true; + } + +private: + Type record_type_; + std::vector logical_fields_; + int record_dimensions_; + + size_t logical_slot(const std::string &name) const; + const StructField &physical_field(const std::string &name) const; +}; + +namespace Internal { +std::vector little_endian_scalar_encode(Type word_type, const std::vector &inputs); +std::vector little_endian_scalar_decode(Type word_type, const std::vector &encoded); +ApproximationSignature little_endian_scalar_signature(Type word_type, const ApproximationPorts &inputs); +} // namespace Internal + +/** Convert a scalar integral word per record to/from a leading little-endian + * byte dimension. Decode deliberately uses concat_bits so struct lowering and + * ordinary byte buffers share the same wide-load optimization path. */ +template +struct LittleEndianScalarPack { + std::vector encode(const std::vector &inputs) const { + return Internal::little_endian_scalar_encode(type_of(), inputs); + } + + std::vector decode(const std::vector &encoded) const { + return Internal::little_endian_scalar_decode(type_of(), encoded); + } + + /** A word per record <-> its bytes in a new leading dimension. */ + ApproximationSignature signature(const ApproximationPorts &inputs) const { + return Internal::little_endian_scalar_signature(type_of(), inputs); + } + + /** Every bit of the word is kept. */ + bool lossless() const { + return true; + } +}; + +/** Exact fixed-width planar packing. For `(field_bits, positions)`, one byte + * contains `8/field_bits` planes, each plane spanning `positions` consecutive + * elements. This component applies no recentering and no lookup policy. */ +struct PlanarFieldPack { + PlanarFieldPack(int field_bits, int positions); + + std::vector encode(const std::vector &inputs) const; + std::vector decode(const std::vector &encoded) const; + + /** (element, record) fields <-> (position, record) bytes. The fields + * must fit in `field_bits`: encode masks each one to that width, so + * values outside [0, 2^field_bits - 1] are silently truncated. */ + ApproximationSignature signature() const; + + /** Exact for fields in the declared input range. */ + bool lossless() const { + return true; + } + +private: + int field_bits_, positions_, planes_; +}; + +} // namespace Halide + +#endif diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 1f02535491cb..5ab6f2abafa4 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -61,6 +61,7 @@ target_sources( AlignLoads.h AllocationBoundsInference.h ApplySplit.h + Approximation.h Argument.h AssociativeOpsTable.h Associativity.h @@ -246,6 +247,7 @@ target_sources( AlignLoads.cpp AllocationBoundsInference.cpp ApplySplit.cpp + Approximation.cpp Argument.cpp AssociativeOpsTable.cpp Associativity.cpp diff --git a/src/Func.cpp b/src/Func.cpp index 1f8f3b7d7146..5af566f4290a 100644 --- a/src/Func.cpp +++ b/src/Func.cpp @@ -3,6 +3,7 @@ #include #include #include +#include #include #include #include @@ -13,6 +14,7 @@ #endif #include "ApplySplit.h" +#include "Approximation.h" #include "Argument.h" #include "Associativity.h" #include "Bounds.h" @@ -1266,6 +1268,94 @@ Func Stage::rfactor(const vector> &preserved) { return intm; } +// How many nodes distribute_products() may rewrite before giving up. Multiplying +// out is exponential in the nesting depth of sums, and a reduction body deep +// enough to hit this has no business being split into that many accumulators. +constexpr int distribute_budget = 1024; + +// Multiply out products over sums, so that a sum reduction's increment becomes +// a flat sum of products: (a + b) * c -> a * c + b * c. This exposes terms with +// *different* invariant factors, which a single flattening of the multiply chain +// cannot see -- in (a + b) * c the sum is one opaque leaf. +// +// Subtraction is deliberately left alone (treated as a leaf): distributing it +// would have to push the sign into one side, which does not group and is not +// meaningful for unsigned wraparound. +Expr distribute_products(const Expr &e, int &budget) { + if (--budget < 0) { + return e; + } + if (e.node_type() == IRNodeType::Add) { + auto [a, b] = *as_binary_operands(e); + return Add::make(distribute_products(a, budget), distribute_products(b, budget)); + } + if (const Mul *mul = e.as()) { + Expr a = distribute_products(mul->a, budget); + Expr b = distribute_products(mul->b, budget); + // Recurse on the rewritten form: either side may itself be a product of + // sums that only became visible after this step. + if (a.node_type() == IRNodeType::Add) { + auto [a0, a1] = *as_binary_operands(a); + return distribute_products(Add::make(Mul::make(a0, b), Mul::make(a1, b)), budget); + } + if (b.node_type() == IRNodeType::Add) { + auto [b0, b1] = *as_binary_operands(b); + return distribute_products(Add::make(Mul::make(a, b0), Mul::make(a, b1)), budget); + } + return Mul::make(a, b); + } + return e; +} + +Stage &Stage::distribute() { + user_assert(!definition.is_init()) << "distribute() must be called on an update definition\n"; + + definition.schedule().touched() = true; + + const auto &prover_result = prove_associativity(function.name(), definition.args(), definition.values()); + user_assert(prover_result.associative()) + << "distribute() requires an associative update definition, but the update " + << "definition of " << function.name() << " is not associative.\n"; + + auto is_self_ref = [&](const Expr &e) { + const Call *c = e.as(); + return c && c->name == function.name() && c->call_type == Call::Halide; + }; + + vector values = definition.values(); + bool changed = false; + for (size_t i = 0; i < values.size(); ++i) { + optional law = distributive_law_for(prover_result.pattern.ops[i]); + if (!law || law->outer_op != IRNodeType::Add || law->inner_op != IRNodeType::Mul) { + continue; + } + // Lets may hide the combiner (e.g. an rfactor of this same update + // introduced promise_clamped bindings), so inline them first. + Expr value = substitute_in_all_lets(values[i]); + optional> split = select_binary_operand(value, law->outer_op, is_self_ref); + if (!split) { + continue; + } + int budget = distribute_budget; + Expr expanded = distribute_products(split->second, budget); + user_assert(budget >= 0) + << "distribute() gave up multiplying out the update definition of " + << function.name() << ": it expands to more than " << distribute_budget + << " nodes.\n"; + if (!equal(expanded, split->second)) { + values[i] = Add::make(split->first, expanded); + changed = true; + } + } + + user_assert(changed) + << "distribute() found no product over a sum to multiply out in the update " + << "definition of " << function.name() << ".\n"; + + definition.values() = values; + return *this; +} + FuncVec Stage::hoist_invariants() { user_assert(!definition.is_init()) << "hoist_invariants() must be called on an update definition\n"; @@ -2778,6 +2868,51 @@ Func Func::clone_in(const vector &fs) { return get_wrapper(func, name() + "_clone", fs, true); } +ApproximationResult Func::approximate_by(const Approximation &p, const vector &consumers) { + EncodeResult enc = p.encode({*this}); + user_assert(!enc.encoded.empty()) + << "approximate_by: Approximation::encode(" << name() << ") returned no Funcs\n"; + + DecodeResult dec = p.decode(enc.encoded); + user_assert(dec.decoded.size() == 1) + << "approximate_by: Approximation::decode() must return exactly one Func (the " + << "round-trip replacement), but returned " << dec.decoded.size() << "\n"; + + Func round_trip = dec.decoded[0]; + user_assert(round_trip.dimensions() == dimensions()) + << "approximate_by: decode(encode(" << name() << "))'s result (" << round_trip.name() + << ") has " << round_trip.dimensions() << " dimensions, but " << name() << " has " + << dimensions() << " -- Approximation implementations must reproduce the original " + << "Func's signature exactly\n"; + user_assert(round_trip.types() == types()) + << "approximate_by: decode(encode(" << name() << "))'s result (" << round_trip.name() + << ") has a different type than " << name() << " -- Approximation implementations " + << "must reproduce the original Func's signature exactly\n"; + + for (const Func &g : consumers) { + user_assert(g.name() != name()) + << "approximate_by: " << name() << " cannot be its own consumer\n"; + // Eager and destructive, like Func::rfactor() and the targeted + // form of Func::in(). + g.function().substitute_calls(func, round_trip.function()); + } + + vector intermediates; + std::set seen = {name(), round_trip.name()}; + auto add = [&](const vector &fs) { + for (const Func &g : fs) { + if (seen.insert(g.name()).second) { + intermediates.push_back(g); + } + } + }; + add(enc.encoded); + add(enc.intermediates); + add(dec.intermediates); + return {round_trip, enc.encoded, enc.encoded_ports, intermediates, enc.stage_outputs, dec.stage_outputs, + std::move(enc.trace), std::move(dec.trace)}; +} + Func Func::copy_to_device(DeviceAPI d) { user_assert(defined()) << "copy_to_device on Func " << name() << " with no definition\n"; diff --git a/src/Func.h b/src/Func.h index d0e1d152f34d..ec492b27c20e 100644 --- a/src/Func.h +++ b/src/Func.h @@ -60,6 +60,8 @@ struct VarOrRVar { class ImageParam; class FuncVec; +class Approximation; +struct ApproximationResult; namespace Internal { struct AssociativeOp; @@ -258,6 +260,32 @@ class Stage { */ FuncVec hoist_invariants(); + /** Multiply out products over sums in this update definition's increment, + * so that hoist_invariants() sees a flat sum of terms. Like rfactor(), this + * must be called on an update definition, and it rewrites that definition in + * place. Returns this Stage, so it can be chained. + * + * (a + b) * c becomes a * c + b * c, recursively. Subtraction is left alone. + * + * This is a schedule decision, not a normalization: whether it pays depends + * on what the terms turn out to contain. It pays when the sum hides operands + * with different loop-invariant factors, since each then reduces to a + * factor-free body of its own: + * \code + * f() += (d*q(r) + m) * (e*p(r)); + * \endcode + * multiplies out to d*e * q(r)*p(r) + m*e * p(r), which hoist_invariants() + * turns into two factor-free accumulators. Left alone, the best it could do + * is hoist e and reduce over the sum. + * + * It costs an accumulator per distinct factor, so it is a pessimization when + * the terms share a factor that was already hoistable -- s * (r + 1) is + * better left as one accumulator with factor s than split into two. + * + * It is an error if there is no product over a sum to multiply out. + */ + Stage &distribute(); + /** Schedule the iteration over this stage to be fused with another * stage 's' from outermost loop to a given LoopLevel. 'this' stage will * be computed AFTER 's' in the innermost fused dimension. There should not @@ -1534,6 +1562,18 @@ class Func { Func clone_in(const std::vector &fs); //@} + /** Eagerly and destructively replace every call to this Func inside + * each Func in 'consumers' with a call to the round trip + * p.decode(p.encode(*this)) instead. The substitution happens + * immediately, like Func::rfactor() and the targeted form of + * Func::in(), but substitutes a different computation rather than an + * identity wrapper -- see Approximation.h and + * doc/Approximation.md for the rationale. Because the + * substitution is eager, it can only rewrite Funcs that are already + * fully defined at the point of the call -- there is no counterpart to + * the global Func::in() that also covers Funcs written later. */ + ApproximationResult approximate_by(const Approximation &p, const std::vector &consumers); + /** Declare that this function should be implemented by a call to * halide_buffer_copy with the given target device API. Asserts * that the Func has a pure definition which is a simple call to a diff --git a/test/correctness/CMakeLists.txt b/test/correctness/CMakeLists.txt index a7ca0e1276a2..d09f54754341 100644 --- a/test/correctness/CMakeLists.txt +++ b/test/correctness/CMakeLists.txt @@ -14,6 +14,13 @@ tests( aligned_split_partition.cpp aligned_split_reduction.cpp ambiguous_inline_reductions.cpp + approximate_by.cpp + approximation_components.cpp + approximation_parallel.cpp + approximation_signatures.cpp + approximation_testing.cpp + approximation_trace.cpp + approximation_unit_forms.cpp argmax.cpp arm_cpu_detect.cpp associativity.cpp diff --git a/test/correctness/approximate_by.cpp b/test/correctness/approximate_by.cpp new file mode 100644 index 000000000000..0ec4a245ebad --- /dev/null +++ b/test/correctness/approximate_by.cpp @@ -0,0 +1,294 @@ +#include "Halide.h" +#include +#include + +using namespace Halide; + +namespace { + +constexpr int kBlockSize = 8; + +// A minimal symmetric integer quantizer -- self-contained (no relation to +// any specific real-world format), just enough to exercise: encode() +// returning multiple Funcs plus a genuine scheduling-only intermediate (the +// per-block amax reduction), decode() combining them back into a single +// Func matching the original's signature, and approximate_by()'s eager +// substitution. +struct SymmetricQuantizer { + std::vector encode(const std::vector &inputs) const { + Func f = inputs[0]; + Var x("x"), i("i"); + RDom r(0, kBlockSize, "r"); + + Func amax("amax"); + amax(i) = 0.0f; + amax(i) = max(amax(i), abs(f(i * kBlockSize + r))); + + Func d("d"); + d(i) = amax(i) / 127.0f; + + Func q("q"); + Expr id = select(d(x / kBlockSize) != 0.0f, 1.0f / d(x / kBlockSize), 0.0f); + q(x) = cast(clamp(round(f(x) * id), -127, 127)); + + return {q, d}; + } + + std::vector decode(const std::vector &encoded) const { + Func q = encoded[0], d = encoded[1]; + Var x("x"); + Func dequantized("dequantized"); + dequantized(x) = cast(q(x)) * d(x / kBlockSize); + return {dequantized}; + } +}; + +// A single-form unit whose intermediates are only discovered. +struct TwoStep { + std::string tag = ""; + + Func encode(const Func &in) const { + Var x("x"); + Func a("step_a" + tag), b("step_b" + tag); + a(x) = in(x) + 1; + b(x) = a(x) * 2; + return b; + } + Func decode(const Func &in) const { + Var x("x"); + Func out("step_out"); + out(x) = in(x) / 2 - 1; + return out; + } +}; + +int check_discovery() { + Var x("x"); + // f has a producer; neither f nor g may show up as an intermediate. + Func g("disc_g"), f("disc_f"), c("disc_c"); + g(x) = x; + f(x) = g(x) * 2; + c(x) = f(x); + ApproximationResult r = f.approximate_by(Approximation(TwoStep{}), {c}); + std::vector names; + for (const Func &i : r.intermediates) { + names.push_back(i.name()); + if (i.name() == "disc_f" || i.name() == "disc_g" || i.name() == r.replacement.name()) { + printf("Intermediates contain an excluded Func: %s\n", i.name().c_str()); + return 1; + } + } + // encoded (step_b) first, then the discovered step_a. + if (names.size() != 2 || names[0].rfind("step_b", 0) != 0 || names[1].rfind("step_a", 0) != 0) { + printf("Unexpected intermediates for TwoStep\n"); + return 1; + } + for (Func i : r.intermediates) { + i.compute_root(); + } + r.replacement.compute_root(); + Buffer out = c.realize({8}); + for (int i = 0; i < 8; i++) { + if (out(i) != i * 2) { + printf("TwoStep round trip wrong at %d\n", i); + return 1; + } + } + + // Topological order: a producer precedes its consumer. + Approximation two = TwoStep{}; + EncodeResult e = two.encode({f}); + Func in("topo_in"); + in(x) = x; + Approximation composed = Compose{TwoStep{"_inner"}, TwoStep{"_outer"}}; + EncodeResult ce = composed.encode({in}); + auto index_of = [&](const char *n) { + for (size_t i = 0; i < ce.intermediates.size(); i++) { + if (ce.intermediates[i].name() == n) { + return (int)i; + } + } + return -1; + }; + // The Compose reports the inter-stage Func (the inner stage's step_b) and + // both stages' step_a Funcs, but not its own output or input. + if (ce.intermediates.size() != 3 || index_of("step_b_inner") < 0 || index_of("step_a_inner") < 0 || + index_of("step_a_outer") < 0) { + printf("Compose intermediates wrong (%zu)\n", ce.intermediates.size()); + return 1; + } + // Each intermediate must appear after everything it calls. + for (size_t i = 0; i < ce.intermediates.size(); i++) { + for (const auto &[n, callee] : Internal::find_direct_calls(ce.intermediates[i].function())) { + for (size_t j = i; j < ce.intermediates.size(); j++) { + if (ce.intermediates[j].function().same_as(callee)) { + printf("Intermediates are not in topological order\n"); + return 1; + } + } + } + } + if (e.intermediates.size() != 1 || e.intermediates[0].name().rfind("step_a", 0) != 0) { + printf("Unexpected TwoStep encode intermediates\n"); + return 1; + } + // Compose's stage_outputs: two children then the Compose itself. + if (ce.stage_outputs.size() != 3 || ce.stage_outputs[2].intermediates.size() != 3 || + ce.stage_outputs[0].intermediates.size() != 1) { + printf("Compose stage_outputs wrong\n"); + return 1; + } + return 0; +} + +} // namespace + +int main(int argc, char **argv) { + if (check_discovery()) { + return 1; + } + + Var x("x"); + + Func f("f"); + f(x) = sin(cast(x) * 0.1f) * 100.0f; + + // g is rewired by approximate_by() below; h is not, and must keep + // seeing the exact, unquantized f. + Func g("g"); + g(x) = f(x) * 2.0f + 1.0f; + + Func h("h"); + h(x) = f(x) * 3.0f; + + SymmetricQuantizer quant; + ApproximationResult result = f.approximate_by(quant, {g}); + + // Stage tracing preserves opaque identities through nested combinators, + // including repeated component types and Parallel/TrustedInverse ownership. + Approximation first_identity = Identity{}, second_identity = Identity{}; + Approximation trusted_encoder = Identity{}, trusted_decoder = Identity{}; + Approximation applied = Parallel{std::vector{second_identity}}; + Approximation trusted = TrustedInverse(trusted_encoder, trusted_decoder); + Approximation nested = Compose(first_identity, applied, trusted); + + // A copy of a handle is the same stage; a second conversion is not. + Approximation first_copy = first_identity; + if (!first_copy.same_as(first_identity) || first_identity.same_as(second_identity) || + Approximation().defined() || !nested.defined()) { + printf("Approximation handle identity semantics are wrong\n"); + return 1; + } + + Func traced_source("traced_source"), traced_consumer("traced_consumer"); + traced_source(x) = cast(x); + traced_consumer(x) = traced_source(x); + ApproximationResult traced = traced_source.approximate_by(nested, {traced_consumer}); + + auto require_port = [&](const Approximation &stage, const char *label) { + if (!traced.encoded_by(stage).defined() || !traced.decoded_by(stage).defined()) { + printf("Missing encode/decode stage trace for %s\n", label); + return false; + } + return true; + }; + if (!require_port(first_identity, "first repeated Identity") || + !require_port(second_identity, "second repeated Identity") || + !require_port(first_copy, "copy of first Identity") || + !require_port(applied, "Parallel") || + !require_port(trusted, "TrustedInverse") || + !require_port(nested, "outer Compose")) { + return 1; + } + // Looking up a stage that wasn't invoked in a direction is an error, so + // check direction-specificity against the raw records instead. + auto recorded = [](const std::vector &outputs, const Approximation &stage) { + for (const ApproximationStageOutputs &o : outputs) { + if (o.stage.same_as(stage)) { + return true; + } + } + return false; + }; + if (!recorded(traced.encoded_stage_outputs, trusted_encoder) || + recorded(traced.decoded_stage_outputs, trusted_encoder) || + recorded(traced.encoded_stage_outputs, trusted_decoder) || + !recorded(traced.decoded_stage_outputs, trusted_decoder)) { + printf("TrustedInverse did not preserve direction-specific child traces\n"); + return 1; + } + Approximation unused = Identity{}; + if (recorded(traced.encoded_stage_outputs, unused) || + recorded(traced.decoded_stage_outputs, unused)) { + printf("Unused stage unexpectedly recorded\n"); + return 1; + } + if (traced.encoded_by(first_identity, 1).defined() || + traced.decoded_by(first_identity, 1).defined()) { + printf("Out-of-range port unexpectedly resolved\n"); + return 1; + } + + // encoded (q, d) come first, then amax, discovered without being declared. + { + std::vector names; + for (const Func &i : result.intermediates) { + names.push_back(i.name()); + } + if (names != std::vector{"q", "d", "amax"}) { + printf("Unexpected intermediates (%zu)\n", names.size()); + return 1; + } + // Stage-local intermediates are reported per stage, too. + const ApproximationStageOutputs &enc_stage = result.encoded_stage_outputs.back(); + if (enc_stage.intermediates.size() != 1 || enc_stage.intermediates[0].name() != "amax") { + printf("Unexpected per-stage intermediates\n"); + return 1; + } + } + result.replacement.compute_root(); + for (Func intermediate : result.intermediates) { + intermediate.compute_root(); + } + + const int kSize = 64; + Buffer g_out = g.realize({kSize}); + Buffer h_out = h.realize({kSize}); + + for (int i = 0; i < kSize; i++) { + const float fx = sinf(i * 0.1f) * 100.0f; + + // Independently recompute the same per-block quantization encode() + // performs, to build a bit-exact reference for what g should see. + const int block = i / kBlockSize; + float amax = 0.0f; + for (int j = 0; j < kBlockSize; j++) { + const float v = sinf((block * kBlockSize + j) * 0.1f) * 100.0f; + amax = std::max(amax, std::fabs(v)); + } + const float d = amax / 127.0f; + const float id = d != 0.0f ? 1.0f / d : 0.0f; + float q = std::round(fx * id); + q = std::max(-127.0f, std::min(127.0f, q)); + const float dequantized = q * d; + + const float expected_g = dequantized * 2.0f + 1.0f; + if (std::fabs(g_out(i) - expected_g) > 1e-5f * std::max(1.0f, std::fabs(expected_g))) { + printf("g(%d) = %f, expected %f -- approximate_by's substitution did not take effect\n", + i, g_out(i), expected_g); + return 1; + } + + // h was never passed as a consumer to approximate_by(): it must + // see the real f, not the quantized round trip. + const float expected_h = fx * 3.0f; + if (std::fabs(h_out(i) - expected_h) > 1e-5f * std::max(1.0f, std::fabs(expected_h))) { + printf("h(%d) = %f, expected %f -- approximate_by affected a Func not in `consumers`\n", + i, h_out(i), expected_h); + return 1; + } + } + + printf("Success!\n"); + return 0; +} diff --git a/test/correctness/approximation_components.cpp b/test/correctness/approximation_components.cpp new file mode 100644 index 000000000000..65c6b11f50c2 --- /dev/null +++ b/test/correctness/approximation_components.cpp @@ -0,0 +1,283 @@ +#include "Halide.h" + +#include +#include + +using namespace Halide; + +namespace { + +// A per-block absmax int8 quantizer on (within, block) floats: codes in +// [-qmax, qmax] and one float scale per block. +struct AbsMaxQuantizer { + int block, qmax; + + std::vector encode(const std::vector &in) const { + Var kk("kk"), blk("blk"); + RDom r(0, block); + Func amax("absmax_stat"), scale("absmax_scale"), codes("absmax_codes"); + amax(blk) = 0.0f; + amax(blk) = max(amax(blk), abs(in[0](r, blk))); + scale(blk) = amax(blk) / (float)qmax; + codes(kk, blk) = cast(round(in[0](kk, blk) / select(scale(blk) == 0.0f, 1.0f, scale(blk)))); + return {codes, scale}; + } + + std::vector decode(const std::vector &encoded) const { + Var kk("kk"), blk("blk"); + Func out("absmax_decoded"); + out(kk, blk) = cast(encoded[0](kk, blk)) * encoded[1](blk); + return {out}; + } + + ApproximationSignature signature() const { + return {{{"block", Float(32), 2}}, {{"codes", Int(8), 2}, {"scale", Float(32), 1}}}; + } +}; + +// Reduced-precision storage: a cast to and from float16. +Pointwise f16_storage() { + return Pointwise{"f16", + [](Expr x) { return strict_float(cast(x)); }, + [](Expr x) { return strict_float(cast(x)); }}; +} + +int test_struct_layout_1d() { + Type record_type = Type::Struct({{"d", Float(16)}, {"qh", UInt(8), 4}, {"qs", UInt(8), 16}}); + Var element("element"), record("record"); + Func qs("qs"), qh("qh"), d("d"); + qs(element, record) = cast(element + 3 * record); + qh(element, record) = cast(0x80 + element + record); + d(record) = cast(cast(record) + 0.5f); + + Approximation layout = StructLayout(record_type, {"qs", "qh", "d"}); + DecodeResult decoded = layout.decode(layout.encode({qs, qh, d}).encoded); + Buffer out_qs = decoded.decoded[0].realize({16, 3}); + Buffer out_qh = decoded.decoded[1].realize({4, 3}); + Buffer out_d = decoded.decoded[2].realize({3}); + for (int r = 0; r < 3; ++r) { + if ((float)out_d(r) != r + 0.5f) { + return 1; + } + for (int i = 0; i < 16; ++i) { + if (out_qs(i, r) != (uint8_t)(i + 3 * r)) { + return 1; + } + } + for (int i = 0; i < 4; ++i) { + if (out_qh(i, r) != (uint8_t)(0x80 + i + r)) { + return 1; + } + } + } + return 0; +} + +int test_struct_layout_2d() { + Type record_type = Type::Struct({{"tag", UInt(16)}, {"pixels", UInt(8), 3}}); + Var element("element"), x("x"), y("y"); + Func pixels("pixels"), tag("tag"); + pixels(element, x, y) = cast(element + 10 * x + 30 * y); + tag(x, y) = cast(100 + x + 4 * y); + Approximation layout = StructLayout(record_type, {"pixels", "tag"}, 2); + DecodeResult decoded = layout.decode(layout.encode({pixels, tag}).encoded); + Buffer out_pixels = decoded.decoded[0].realize({3, 4, 2}); + Buffer out_tag = decoded.decoded[1].realize({4, 2}); + for (int yy = 0; yy < 2; ++yy) { + for (int xx = 0; xx < 4; ++xx) { + if (out_tag(xx, yy) != 100 + xx + 4 * yy) { + return 1; + } + for (int i = 0; i < 3; ++i) { + if (out_pixels(i, xx, yy) != i + 10 * xx + 30 * yy) { + return 1; + } + } + } + } + return 0; +} + +int test_struct_layout_contract_errors() { + if (!Halide::exceptions_enabled()) { + return 0; + } + Type record_type = Type::Struct({{"tag", UInt(16)}, {"pixels", UInt(8), 3}}); + Var element("element"), record("record"); + Func pixels("pixels"), wrong_tag("wrong_tag"); + pixels(element, record) = cast(element); + wrong_tag(record) = cast(record); + try { + Approximation layout = StructLayout(record_type, {"pixels", "tag"}); + (void)layout.encode({pixels, wrong_tag}); + return 1; + } catch (const CompileError &) { + } + try { + StructLayout duplicate(record_type, {"pixels", "pixels"}); + return 1; + } catch (const CompileError &) { + } + return 0; +} + +int test_scalar_components() { + Var record("record"); + Func values("values"); + values(record) = cast(record) / 3.0f; + Approximation storage = f16_storage(); + EncodeResult stored = storage.encode({values}); + DecodeResult cast_roundtrip = storage.decode(stored.encoded); + Buffer out = cast_roundtrip.decoded[0].realize({8}); + for (int i = 0; i < 8; ++i) { + float expected = (float)(float16_t)(i / 3.0f); + if (out(i) != expected) { + printf("f16 cast mismatch at %d: %g vs %g\n", i, out(i), expected); + return 1; + } + } + + Func words("words"); + words(record) = cast((int32_t)0x10203040) + cast(record); + Approximation little_endian = LittleEndianScalarPack{}; + EncodeResult bytes = little_endian.encode({words}); + DecodeResult word_roundtrip = little_endian.decode(bytes.encoded); + Buffer packed = bytes.encoded[0].realize({4, 5}); + Buffer unpacked = word_roundtrip.decoded[0].realize({5}); + for (int r = 0; r < 5; ++r) { + if (unpacked(r) != 0x10203040u + r || packed(0, r) != (uint8_t)(0x40 + r) || + packed(1, r) != 0x30 || packed(2, r) != 0x20 || packed(3, r) != 0x10) { + printf("LittleEndian mismatch at %d: %08x [%02x %02x %02x %02x]\n", r, + unpacked(r), packed(0, r), packed(1, r), packed(2, r), packed(3, r)); + return 1; + } + } + return 0; +} + +int test_code_components() { + Var element("element"), record("record"); + Func signed_nibbles("signed_nibbles"); + signed_nibbles(element, record) = cast((element % 16) - 8); + Approximation offset = Pointwise{"offset", + [](Expr x) { return cast(x + 8); }, + [](Expr x) { return cast(x - 8); }}; + EncodeResult offset_codes = offset.encode({signed_nibbles}); + DecodeResult signed_roundtrip = offset.decode(offset_codes.encoded); + Buffer stored_nibbles = offset_codes.encoded[0].realize({16, 2}); + Buffer restored_nibbles = signed_roundtrip.decoded[0].realize({16, 2}); + for (int r = 0; r < 2; ++r) { + for (int i = 0; i < 16; ++i) { + if (stored_nibbles(i, r) != i || restored_nibbles(i, r) != i - 8) { + return 1; + } + } + } + + Func nibbles("nibbles"); + nibbles(element, record) = cast(element % 16); + Approximation planar = PlanarFieldPack(4, 16); + EncodeResult planar_bytes = planar.encode({nibbles}); + planar_bytes.encoded[0].compute_root(); + DecodeResult planar_fields = planar.decode(planar_bytes.encoded); + Buffer bytes_out = planar_bytes.encoded[0].realize({16, 1}); + Buffer fields_out = planar_fields.decoded[0].realize({32, 1}); + for (int i = 0; i < 16; ++i) { + if (bytes_out(i, 0) != (uint8_t)(i | (i << 4)) || + fields_out(i, 0) != i || fields_out(i + 16, 0) != i) { + return 1; + } + } + return 0; +} + +int test_block_components() { + Var k("k"); + Func flat("flat"); + flat(k) = cast(k); + Approximation reshape = BlockReshape(32); + DecodeResult reshaped = reshape.decode(reshape.encode({flat}).encoded); + Buffer roundtrip = reshaped.decoded[0].realize({96}); + for (int i = 0; i < 96; ++i) { + if (roundtrip(i) != i) { + return 1; + } + } + + return 0; +} + +int test_standard_quant_compositions() { + Var k("k"); + + Pointwise offset{"offset", + [](Expr x) { return cast(x + 8); }, + [](Expr x) { return cast(x - 8); }}; + Type q4_type = Type::Struct({{"d", Float(16)}, {"qs", UInt(8), 16}}); + Func q4_values("q4_values"); + q4_values(k) = cast((k % 15) - 7); + Approximation q4 = Compose( + BlockReshape{32}, + AbsMaxQuantizer{32, 7}, + Parallel{{"codes", Compose{offset, PlanarFieldPack{4, 16}}}, + {"scale", f16_storage()}}, + StructLayout{q4_type, {"qs", "d"}}); + EncodeResult q4_encoded = q4.encode({q4_values}); + for (Func intermediate : q4_encoded.intermediates) { + intermediate.compute_root(); + } + DecodeResult q4_decoded = q4.decode(q4_encoded.encoded); + Buffer q4_roundtrip = q4_decoded.decoded[0].realize({64}); + for (int i = 0; i < 64; ++i) { + if (q4_roundtrip(i) != (i % 15) - 7) { + return 1; + } + } + + Type q8_type = Type::Struct({{"d", Float(16)}, {"qs", Int(8), 32}}); + Func q8_values("q8_values"); + Expr local = k % 32; + q8_values(k) = cast(select(local == 0, -127, local - 16)); + Approximation q8 = Compose( + BlockReshape{32}, + AbsMaxQuantizer{32, 127}, + Parallel{Identity{}, f16_storage()}, + StructLayout{q8_type, {"qs", "d"}}); + EncodeResult q8_encoded = q8.encode({q8_values}); + for (Func intermediate : q8_encoded.intermediates) { + intermediate.compute_root(); + } + DecodeResult q8_decoded = q8.decode(q8_encoded.encoded); + Buffer q8_roundtrip = q8_decoded.decoded[0].realize({64}); + for (int i = 0; i < 64; ++i) { + int local_i = i % 32; + float expected = local_i == 0 ? -127.0f : local_i - 16.0f; + if (q8_roundtrip(i) != expected) { + return 1; + } + } + return 0; +} + +} // namespace + +int main(int argc, char **argv) { + struct Test { + const char *name; + int (*run)(); + } tests[] = {{"StructLayout 1-D", test_struct_layout_1d}, + {"StructLayout 2-D", test_struct_layout_2d}, + {"StructLayout contract errors", test_struct_layout_contract_errors}, + {"scalar packs", test_scalar_components}, + {"code packs", test_code_components}, + {"block components", test_block_components}, + {"standard quant compositions", test_standard_quant_compositions}}; + for (const Test &test : tests) { + if (test.run()) { + printf("Approximation component test failed: %s\n", test.name); + return 1; + } + } + printf("Success!\n"); + return 0; +} diff --git a/test/correctness/approximation_parallel.cpp b/test/correctness/approximation_parallel.cpp new file mode 100644 index 000000000000..ce95942d64be --- /dev/null +++ b/test/correctness/approximation_parallel.cpp @@ -0,0 +1,250 @@ +#include "Halide.h" +#include +#include + +using namespace Halide; + +namespace { + +#define CHECK(cond) \ + do { \ + if (!(cond)) { \ + printf("%s:%d: check failed: %s\n", __FILE__, __LINE__, #cond); \ + return 1; \ + } \ + } while (0) + +using Names = std::vector; + +Names names(const ApproximationPorts &ports) { + Names n; + for (const ApproximationPort &p : ports) { + n.push_back(p.name); + } + return n; +} + +// values -> (codes, scale), with a fixed scale of one half. +struct Producer { + std::vector encode(const std::vector &in) const { + Var x("x"); + Func codes("codes"), scale("scale"); + scale(x) = 0.5f; + codes(x) = cast(round(in[0](x) / scale(x))); + return {codes, scale}; + } + std::vector decode(const std::vector &in) const { + Var x("x"); + Func values("values"); + values(x) = cast(in[0](x)) * in[1](x); + return {values}; + } + ApproximationSignature signature() const { + return {{{"values", Float(32), 1}}, + {{"codes", Int(8), 1}, {"scale", Float(32), 1}}}; + } +}; + +// One signed byte to a low and a high nibble: one port becomes two. +struct Nibbles { + std::vector encode(const std::vector &in) const { + Var x("x"); + Func lo("lo_nibble"), hi("hi_nibble"); + lo(x) = in[0](x) & cast(15); + hi(x) = in[0](x) >> 4; + return {lo, hi}; + } + std::vector decode(const std::vector &in) const { + Var x("x"); + Func v("joined"); + v(x) = cast(in[1](x) * 16 + in[0](x)); + return {v}; + } + ApproximationSignature signature() const { + return {{{"byte", Int(8), 1}}, {{"lo", Int(8), 1}, {"hi", Int(8), 1}}}; + } + bool lossless() const { + return true; + } +}; + +// A cast to and from float16. +Pointwise f16_storage() { + return Pointwise{"f16", + [](Expr x) { return strict_float(cast(x)); }, + [](Expr x) { return strict_float(cast(x)); }} + .with_types(Float(32), Float(16)); +} + +// (within, block) floats -> int8 codes and a float scale per block. +struct AbsMaxQuantizer { + int block, qmax; + std::vector encode(const std::vector &in) const { + Var kk("kk"), blk("blk"); + RDom r(0, block); + Func amax("absmax_stat"), scale("absmax_scale"), codes("absmax_codes"); + amax(blk) = 0.0f; + amax(blk) = max(amax(blk), abs(in[0](r, blk))); + scale(blk) = amax(blk) / (float)qmax; + codes(kk, blk) = cast(round(in[0](kk, blk) / select(scale(blk) == 0.0f, 1.0f, scale(blk)))); + return {codes, scale}; + } + std::vector decode(const std::vector &encoded) const { + Var kk("kk"), blk("blk"); + Func out("absmax_decoded"); + out(kk, blk) = cast(encoded[0](kk, blk)) * encoded[1](blk); + return {out}; + } + ApproximationSignature signature() const { + return {{{"block", Float(32), 2}}, {{"codes", Int(8), 2}, {"scale", Float(32), 1}}}; + } +}; + +Func source() { + Var x("x"); + Func f("src"); + f(x) = (cast(x) - 20.0f) * 0.5f; + return f; +} + +int check_round_trip(const Approximation &a, const std::vector &encoded_from_encode, + const EncodeResult &e, Func f) { + for (Func i : e.intermediates) { + i.compute_root(); + } + DecodeResult d = a.decode(e.encoded); + Buffer out = d.decoded[0].realize({40}); + Buffer ref = f.realize({40}); + for (int i = 0; i < 40; i++) { + CHECK(out(i) == ref(i)); + } + (void)encoded_from_encode; + return 0; +} + +int test_named() { + Func f = source(); + Approximation scheme = Compose{Producer{}, + Parallel{{"codes", Nibbles{}}, + {"scale", f16_storage()}}}; + EncodeResult e = scheme.encode({f}); + CHECK(names(e.encoded_ports) == Names({"lo", "hi", "scale"})); + CHECK(e.encoded.size() == 3); + DecodeResult d = scheme.decode(e.encoded); + CHECK(names(d.decoded_ports) == Names({"values"})); + + // The children of the Parallel see the names that flow in. + const ApproximationTraceNode &par = e.trace.children[1]; + CHECK(par.input_names == Names({"codes", "scale"})); + CHECK(par.port_names == Names({"lo", "hi", "scale"})); + const ApproximationTraceNode &dpar = d.trace.children[0]; + CHECK(dpar.input_names == Names({"lo", "hi", "scale"})); + CHECK(dpar.port_names == Names({"codes", "scale"})); + + Approximation alone = scheme; + ApproximationSignature sig = alone.signature({{"values"}}); + CHECK(names(sig.outputs) == Names({"lo", "hi", "scale"})); + + // Ports the Parallel does not mention pass through unchanged. + Approximation only_codes = Compose{Producer{}, Parallel{{"codes", Nibbles{}}}}; + EncodeResult oe = only_codes.encode({f}); + CHECK(names(oe.encoded_ports) == Names({"lo", "hi", "scale"})); + + // A name that flows in wins over a child's declared one. + Approximation renamed = Compose{Producer{}, Parallel{{"codes", Identity{}}}}; + CHECK(names(renamed.encode({f}, {{"weights"}}).encoded_ports) == Names({"codes", "scale"})); + + // Values round trip (fp16 represents multiples of one half exactly). + return check_round_trip(scheme, e.encoded, e, f); +} + +int test_positional() { + Func f = source(); + Approximation scheme = Compose{Producer{}, Parallel{Nibbles{}, f16_storage()}}; + EncodeResult e = scheme.encode({f}); + CHECK(names(e.encoded_ports) == Names({"lo", "hi", "scale"})); + CHECK(names(scheme.decode(e.encoded).decoded_ports) == Names({"values"})); + CHECK(check_round_trip(scheme, e.encoded, e, f) == 0); + + // Identity passes one Func through; widths come from the signatures. + Approximation with_identity = Compose{Producer{}, Parallel{Identity{}, f16_storage()}}; + EncodeResult ie = with_identity.encode({f}); + CHECK(names(ie.encoded_ports) == Names({"codes", "scale"})); + CHECK(check_round_trip(with_identity, ie.encoded, ie, f) == 0); + return 0; +} + +template +bool fails_with(F &&fn, const std::string &needle) { + try { + fn(); + } catch (const CompileError &e) { + return std::string(e.what()).find(needle) != std::string::npos; + } + return false; +} + +int test_errors() { + if (!Halide::exceptions_enabled()) { + return 0; + } + Func f = source(), g = source(); + Approximation missing = Parallel{{"nope", Identity{}}}; + CHECK(fails_with([&] { missing.encode({f, g}, {{"a"}, {"b"}}); }, "nope")); + Approximation ambiguous = Parallel{{"a", Identity{}}}; + CHECK(fails_with([&] { ambiguous.encode({f, g}, {{"a"}, {"a"}}); }, "a")); + Approximation twice = Parallel{{"a", Identity{}}, {"a", Identity{}}}; + CHECK(fails_with([&] { twice.encode({f, g}, {{"a"}, {"b"}}); }, "a")); + Approximation too_few = Parallel{Identity{}, Identity{}}; + CHECK(fails_with([&] { too_few.encode({f}); }, "2")); + CHECK(fails_with([&] { too_few.encode({f, g, f}); }, "3")); + return 0; +} + +// A Q4_0-shaped scheme: BlockReshape, quantize, Parallel, StructLayout. +int test_mirror_and_uniqueness() { + Type block = Type::Struct({{"d", Float(16)}, {"qs", UInt(8), 16}}); + Pointwise offset{"offset", + [](Expr x) { return cast(x + 8); }, + [](Expr x) { return cast(x - 8); }}; + Approximation scheme = Compose{ + BlockReshape{32}, + AbsMaxQuantizer{32, 7}, + Parallel{{"codes", Compose{offset, PlanarFieldPack{4, 16}}}, + {"scale", f16_storage()}}, + StructLayout{block, {"qs", "d"}}}; + Var k("k"); + Func f("q4_src"); + f(k) = cast((k % 15) - 7); + EncodeResult e = scheme.encode({f}); + DecodeResult d = scheme.decode(e.encoded); + CHECK(names(d.decoded_ports) == Names({"values"})); + + // A stand-alone decode agrees with the encode's context on every stage. + const size_t n = e.trace.children.size(); + CHECK(n == 4 && d.trace.children.size() == n); + for (size_t i = 0; i < n; i++) { + const ApproximationTraceNode &en = e.trace.children[i]; + const ApproximationTraceNode &de = d.trace.children[n - 1 - i]; + CHECK(en.input_names == de.port_names); + CHECK(en.port_names == de.input_names); + CHECK(std::set(de.port_names.begin(), de.port_names.end()).size() == de.port_names.size()); + } + std::vector outs = d.stage_outputs; + for (const ApproximationStageOutputs &o : outs) { + CHECK(std::set(o.port_names.begin(), o.port_names.end()).size() == o.port_names.size()); + } + // The struct layout maps slots positionally: flowing names survive it. + CHECK(e.trace.children[3].input_names == Names({"bytes", "scale"})); + return 0; +} + +} // namespace + +int main() { + if (test_named() || test_positional() || test_errors() || test_mirror_and_uniqueness()) { + return 1; + } + printf("Success!\n"); + return 0; +} diff --git a/test/correctness/approximation_signatures.cpp b/test/correctness/approximation_signatures.cpp new file mode 100644 index 000000000000..5e6af77cf06e --- /dev/null +++ b/test/correctness/approximation_signatures.cpp @@ -0,0 +1,357 @@ +#include "Halide.h" +#include +#include +#include + +using namespace Halide; + +namespace { + +#define CHECK(cond) \ + do { \ + if (!(cond)) { \ + printf("%s:%d: check failed: %s\n", __FILE__, __LINE__, #cond); \ + return 1; \ + } \ + } while (0) + +// Declares its ports statically. values -> (codes, scale), one scale per value. +struct Producer { + std::vector encode(const std::vector &in) const { + Var x("x"); + Func codes("codes"), scale("scale"); + scale(x) = 0.5f; + codes(x) = cast(round(in[0](x) / scale(x))); + return {codes, scale}; + } + std::vector decode(const std::vector &in) const { + Var x("x"); + Func values("values"); + values(x) = cast(in[0](x)) * in[1](x); + return {values}; + } + ApproximationSignature signature() const { + return {{{"values", Float(32), 1}}, + {{"codes", Int(8), 1}, {"scale", Float(32), 1}}}; + } +}; + +// Float to float16 bits and back, with no declaration. +struct HalfBits { + Func encode(const Func &f) const { + Func g("half_bits"); + g(_) = reinterpret(cast(f(_))); + return g; + } + Func decode(const Func &f) const { + Func g("half_value"); + g(_) = cast(reinterpret(f(_))); + return g; + } +}; + +// The same, declared. +struct DeclaredHalfBits : HalfBits { + ApproximationSignature signature() const { + return {{{"value", Float(32), 1}}, {{"bits", UInt(16), 1}}}; + } +}; + +// One input, two outputs, undeclared: the outputs are named positionally. +struct SplitParity { + std::vector encode(const std::vector &in) const { + Var x("x"); + Func even("even"), odd("odd"); + even(x) = in[0](2 * x); + odd(x) = in[0](2 * x + 1); + return {even, odd}; + } + std::vector decode(const std::vector &in) const { + Var x("x"); + Func out("merged"); + out(x) = select(x % 2 == 0, in[0](x / 2), in[1](x / 2)); + return {out}; + } +}; + +// One input, two float outputs, declared: values -> (lo, hi). +struct LoHi { + std::vector encode(const std::vector &in) const { + Var x("x"); + Func lo("lo_part"), hi("hi_part"); + lo(x) = in[0](2 * x); + hi(x) = in[0](2 * x + 1); + return {lo, hi}; + } + std::vector decode(const std::vector &in) const { + Var x("x"); + Func out("lohi_merged"); + out(x) = select(x % 2 == 0, in[0](x / 2), in[1](x / 2)); + return {out}; + } + ApproximationSignature signature() const { + return {{{"values", Float(32), 1}}, {{"lo", Float(32), 1}, {"hi", Float(32), 1}}}; + } +}; + +// Records the ports it is handed, and passes everything through. +struct Probe { + std::shared_ptr seen = std::make_shared(); + + std::vector encode(const std::vector &in, const ApproximationPorts &ports) const { + *seen = ports; + return in; + } + std::vector decode(const std::vector &in, const ApproximationPorts &) const { + return in; + } +}; + +struct DeclaredProbe : Probe { + ApproximationSignature signature() const { + return {{{"values"}}, {{"values"}}}; + } +}; + +std::vector names(const ApproximationPorts &ports) { + std::vector result; + for (const ApproximationPort &p : ports) { + result.push_back(p.name); + } + return result; +} + +using Names = std::vector; + +} // namespace + +int main() { + // The constructors accept the natural spellings. + { + ApproximationPort a{"scale"}; + ApproximationPort b{"scale", Float(32)}; + ApproximationPort c{"scale", Float(32), 1}; + ApproximationPorts ports{{"a"}, {"b", UInt(8)}, {"c", Int(8), 2}}; + CHECK(!a.type && !a.dimensions); + CHECK(*b.type == Float(32) && !b.dimensions); + CHECK(*c.dimensions == 1); + CHECK(names(ports) == Names({"a", "b", "c"})); + } + + Var x("x"); + Func f("f"); + f(x) = cast(x % 20) * 0.5f; + Func consumer("consumer"); + consumer(x) = f(x) + 0.0f; + + // A named Parallel in a Compose: half-precision scales, found by name. + Approximation producer = Producer{}; + Approximation half = HalfBits{}; + Approximation apply = Parallel{{"scale", half}}; + Approximation scheme = Compose{producer, apply}; + { + ApproximationResult r = f.approximate_by(scheme, {consumer}); + Buffer out = consumer.realize({64}); + for (int i = 0; i < 64; i++) { + CHECK(out(i) == (float)(i % 20) * 0.5f); + } + + CHECK(r.encoded.size() == 2); + CHECK(names(r.encoded_ports) == Names({"codes", "scale"})); + CHECK(*r.encoded_ports[0].type == Int(8)); + CHECK(*r.encoded_ports[1].type == UInt(16)); + CHECK(r.encoded[1].types()[0] == UInt(16)); + + CHECK(r.encoded_by(producer, "codes").name() == r.encoded[0].name()); + CHECK(r.encoded_by(producer, "scale").types()[0] == Float(32)); + CHECK(r.encoded_by(apply, "scale").name() == r.encoded[1].name()); + CHECK(r.encoded_by(half, "scale").name() == r.encoded[1].name()); + CHECK(r.decoded_by(half, "scale").types()[0] == Float(32)); + CHECK(r.decoded_by(apply, "scale").name() == r.decoded_by(half, "scale").name()); + CHECK(r.decoded_by(apply, "codes").name() == r.encoded[0].name()); + CHECK(r.decoded_by(producer, "values").name() == r.replacement.name()); + // The index forms are unchanged, and a literal 0 is still an index. + CHECK(r.encoded_by(producer, 1).name() == r.encoded_by(producer, "scale").name()); + CHECK(r.encoded_by(producer, 0).name() == r.encoded_by(producer, "codes").name()); + CHECK(r.encoded_by(producer).name() == r.encoded_by(producer, "codes").name()); + + std::ostringstream trace; + trace << r.encode_trace; + CHECK(trace.str().find("scale=half_bits") != std::string::npos); + } + + // Pass-through naming through undeclared single-form units, at any arity. + { + Approximation neg = Pointwise{"neg", [](Expr e) { return -e; }, [](Expr e) { return -e; }}; + EncodeResult e = neg.encode({f}, {ApproximationPort("weights")}); + CHECK(names(e.encoded_ports) == Names({"weights"})); + CHECK(e.trace.port_names == Names({"weights"})); + DecodeResult d = neg.decode(e.encoded, {ApproximationPort("weights")}); + CHECK(names(d.decoded_ports) == Names({"weights"})); + CHECK(d.trace.input_names == Names({"weights"})); + + Approximation chain = Compose{neg, half}; + EncodeResult ce = chain.encode({f}, {ApproximationPort("weights")}); + CHECK(names(ce.encoded_ports) == Names({"weights"})); + + // Without ports, a unit with no declaration names its input positionally. + CHECK(names(neg.encode({f}).encoded_ports) == Names({"0"})); + + // approximate_by gives the root no names, so it uses its declared + // defaults, or else positional ones. + Probe undeclared; + Approximation p1 = undeclared; + (void)f.approximate_by(p1, {}); + CHECK(names(*undeclared.seen) == Names({"0"})); + DeclaredProbe declared; + Approximation p2 = declared; + (void)f.approximate_by(p2, {}); + CHECK(names(*declared.seen) == Names({"values"})); + } + + // Decode on its own: the declared outputs name the encoded inputs. + { + EncodeResult e = scheme.encode({f}); + CHECK(names(e.encoded_ports) == Names({"codes", "scale"})); + // Sever the ports, as sever does. + DecodeResult d = scheme.decode(e.encoded); + CHECK(d.decoded.size() == 1); + CHECK(names(d.decoded_ports) == Names({"values"})); + CHECK(d.trace.children.size() == 2); + CHECK(d.trace.children[0].port_names == Names({"codes", "scale"})); + + DecodeResult pd = producer.decode(producer.encode({f}).encoded); + CHECK(names(pd.decoded_ports) == Names({"values"})); + } + + // Resolved signatures. + Approximation declared_half = DeclaredHalfBits{}; + Approximation declared_scheme = Compose{producer, Parallel{{"scale", declared_half}}}; + { + ApproximationSignature s = declared_scheme.signature(); + CHECK(s.known); + CHECK(names(s.inputs) == Names({"values"})); + CHECK(*s.inputs[0].type == Float(32) && *s.inputs[0].dimensions == 1); + CHECK(names(s.outputs) == Names({"codes", "bits"})); + CHECK(*s.outputs[0].type == Int(8)); + CHECK(*s.outputs[1].type == UInt(16)); + + // Declared input names are only defaults: names that flow in win, for + // static and contextual signatures alike. + CHECK(names(producer.signature().inputs) == Names({"values"})); + CHECK(names(producer.signature({ApproximationPort("other")}).inputs) == Names({"other"})); + CHECK(*producer.signature({ApproximationPort("other")}).inputs[0].type == Float(32)); + ApproximationSignature converted = Approximation(Pointwise{"cast", + [](Expr x) { return cast(x); }, + [](Expr x) { return cast(x); }} + .with_types(Float(32), UInt(16))) + .signature({ApproximationPort("w", std::nullopt, 3)}); + CHECK(names(converted.inputs) == Names({"w"}) && names(converted.outputs) == Names({"w"})); + CHECK(*converted.inputs[0].type == Float(32) && *converted.outputs[0].type == UInt(16)); + CHECK(*converted.outputs[0].dimensions == 3); + + // Undeclared: a single-form unit is pass-through; a multi-form one is unknown. + ApproximationSignature single = half.signature({ApproximationPort("v")}); + CHECK(single.known && names(single.outputs) == Names({"v"})); + CHECK(!Approximation(SplitParity{}).signature().known); + CHECK(!Approximation(Compose{half, SplitParity{}}).signature().known); + } + + // describe() renders the tree without running anything. + { + const char *expected = + "Compose (values: float32 x1) -> (codes: int8 x1, bits: uint16 x1)\n" + " Producer (values: float32 x1) -> (codes: int8 x1, scale: float32 x1)\n" + " Parallel (codes: int8 x1, scale: float32 x1) -> (codes: int8 x1, bits: uint16 x1)\n" + " DeclaredHalfBits (scale: float32 x1) -> (bits: uint16 x1)\n"; + std::string got = declared_scheme.describe(); + if (got != expected) { + printf("Unexpected describe():\n%s", got.c_str()); + return 1; + } + std::ostringstream stream; + stream << declared_scheme; + CHECK(stream.str() == expected); + + std::string undeclared = Approximation(Compose{half, SplitParity{}}).describe(); + const char *expected_undeclared = + "Compose (unknown signature)\n" + " HalfBits (0) -> (0)\n" + " SplitParity (unknown signature)\n"; + if (undeclared != expected_undeclared) { + printf("Unexpected describe():\n%s", undeclared.c_str()); + return 1; + } + } + + // An undeclared multi-form unit changing arity names its outputs positionally. + { + Approximation split = SplitParity{}; + EncodeResult e = split.encode({f}, {ApproximationPort("row")}); + CHECK(names(e.encoded_ports) == Names({"0", "1"})); + DecodeResult d = split.decode(e.encoded); + CHECK(names(d.decoded_ports) == Names({"0"})); + + Func even_source("even_source"); + even_source(x) = x; + ApproximationResult r = even_source.approximate_by(split, {}); + CHECK(r.encoded_by(split, "0").name() == r.encoded[0].name()); + CHECK(r.encoded_by(split, "1").name() == r.encoded[1].name()); + CHECK(r.decoded_by(split, "0").name() == r.replacement.name()); + std::ostringstream trace; + trace << r; + CHECK(trace.str().find("0=even") != std::string::npos); + } + + // Decode's output ports are named like encode's input ports, so a by-name + // Parallel after a Permute works in both directions, with and without an + // encode context. + { + Approximation lohi = LoHi{}; + Approximation lo_half = Parallel{{"lo", half}}; + Approximation permute = Permute{{1, 0}}; + Approximation permuted = Compose{lohi, lo_half, permute}; + + EncodeResult e = permuted.encode({f}); + CHECK(names(e.encoded_ports) == Names({"hi", "lo"})); + CHECK(*e.encoded_ports[1].type == UInt(16)); + + DecodeResult d = permuted.decode(e.encoded); + CHECK(names(d.decoded_ports) == Names({"values"})); + CHECK(d.trace.children.size() == 3); + CHECK(d.decoded[0].types()[0] == Float(32)); + + // Each stage sees the names it had on the way in. + CHECK(d.trace.children[0].input_names == Names({"hi", "lo"})); + CHECK(d.trace.children[0].port_names == Names({"lo", "hi"})); + CHECK(d.trace.children[1].input_names == Names({"lo", "hi"})); + CHECK(d.trace.children[1].port_names == Names({"lo", "hi"})); + CHECK(d.trace.children[2].input_names == Names({"lo", "hi"})); + + // Permute alone, with and without context. + EncodeResult pe = permute.encode({f, consumer}, {{"a"}, {"b"}}); + CHECK(names(pe.encoded_ports) == Names({"b", "a"})); + CHECK(names(permute.decode(pe.encoded, {{"a"}, {"b"}}).decoded_ports) == Names({"a", "b"})); + CHECK(names(permute.decode(pe.encoded).decoded_ports) == Names({"0", "1"})); + Permute rotate{{1, 2, 0}}; + EncodeResult re = Approximation(rotate).encode({f, consumer, f}, {{"a"}, {"b"}, {"c"}}); + CHECK(names(re.encoded_ports) == Names({"b", "c", "a"})); + CHECK(names(Approximation(rotate).decode(re.encoded, {{"a"}, {"b"}, {"c"}}).decoded_ports) == Names({"a", "b", "c"})); + + // Identity and a positional Parallel keep names in decode too. + Approximation ident = Compose{lohi, Parallel{Identity{}, half}, Identity{}}; + EncodeResult ie = ident.encode({f}); + CHECK(names(ie.encoded_ports) == Names({"lo", "hi"})); + CHECK(names(ident.decode(ie.encoded).decoded_ports) == Names({"values"})); + Approximation by_name_first = Compose{lohi, permute, permute, Parallel{{"lo", half}}}; + EncodeResult be = by_name_first.encode({f}); + CHECK(names(by_name_first.decode(be.encoded).decoded_ports) == Names({"values"})); + Approximation two_perms = Compose{lohi, permute, Parallel{{"hi", half}}, permute}; + EncodeResult tpe = two_perms.encode({f}); + CHECK(names(tpe.encoded_ports) == Names({"lo", "hi"})); + CHECK(names(two_perms.decode(tpe.encoded).decoded_ports) == Names({"values"})); + } + + printf("Success!\n"); + return 0; +} diff --git a/test/correctness/approximation_testing.cpp b/test/correctness/approximation_testing.cpp new file mode 100644 index 000000000000..5fbb6f60662a --- /dev/null +++ b/test/correctness/approximation_testing.cpp @@ -0,0 +1,414 @@ +#include "Halide.h" +#include "halide_approximation_testing.h" + +#include +#include + +using namespace Halide; +using namespace Halide::ApproximationTesting; + +namespace { + +#define CHECK(cond) \ + do { \ + if (!(cond)) { \ + printf("%s:%d: check failed: %s\n", __FILE__, __LINE__, #cond); \ + return 1; \ + } \ + } while (0) + +bool same_buffers(const Buffer<> &a, const Buffer<> &b) { + if (a.number_of_elements() != b.number_of_elements() || a.type() != b.type()) { + return false; + } + return memcmp(a.data(), b.data(), a.size_in_bytes()) == 0; +} + +int test_generators() { + // Pinned so that generated data is reproducible across platforms. + uint64_t state = 0; + CHECK(splitmix64(state) == 0xe220a8397b1dcdafULL); + + std::vector dists = { + Distribution::uniform(-1, 1), + Distribution::uniform_int(-5, 5), + Distribution::normal(0, 1), + Distribution::constant(3), + Distribution::zeros(), + Distribution::blockwise_constant(4, Distribution::uniform(0, 1)), + Distribution::outliers(Distribution::normal(0, 1), 8, 100), + Distribution::mixture({Distribution::zeros(), Distribution::normal(0, 1)}, 4), + }; + for (const Distribution &d : dists) { + Buffer<> a = generate(d, Float(32), {32, 3}, 7); + Buffer<> b = generate(d, Float(32), {32, 3}, 7); + CHECK(same_buffers(a, b)); + CHECK(a.dimensions() == 2 && a.dim(0).extent() == 32 && a.dim(1).extent() == 3); + } + // Only the random ones differ across seeds. + for (int i : {0, 1, 2, 5, 6, 7}) { + CHECK(!same_buffers(generate(dists[i], Float(32), {32, 3}, 1), generate(dists[i], Float(32), {32, 3}, 2))); + } + + // Ranges and structure. + Buffer u = generate(Distribution::uniform(2, 3), Float(32), {200}, 1); + for (int i = 0; i < 200; i++) { + CHECK(u(i) >= 2 && u(i) <= 3); + } + Buffer ui = generate(Distribution::uniform_int(-3, 4), Int(8), {200}, 1); + for (int i = 0; i < 200; i++) { + CHECK(ui(i) >= -3 && ui(i) <= 4); + } + Buffer bc = generate(Distribution::blockwise_constant(4, Distribution::uniform(0, 1)), Float(32), {16}, 3); + for (int i = 0; i < 16; i++) { + CHECK(bc(i) == bc(i / 4 * 4)); + } + Buffer ol = generate(Distribution::outliers(Distribution::uniform(-1, 1), 8, 100), Float(32), {16}, 3); + for (int b = 0; b < 2; b++) { + int big = 0; + for (int i = 0; i < 8; i++) { + big += std::abs(ol(b * 8 + i)) >= 50; + } + CHECK(big == 1); + } + Buffer ex = generate(Distribution::extremes(UInt(8)), UInt(8), {64}, 1); + for (int i = 0; i < 64; i++) { + CHECK(ex(i) == 0 || ex(i) == 255); + } + return 0; +} + +// A per-block absmax quantizer on (within, block) floats: codes in +// [-qmax, qmax] and one float scale per block. Rounding to nearest keeps the +// error within half a scale step (plus slack for float rounding). +struct AbsMaxQuantizer { + int block, qmax; + + std::vector encode(const std::vector &in) const { + Var kk("kk"), blk("blk"); + RDom r(0, block); + Func amax("absmax_stat"), scale("absmax_scale"), codes("absmax_codes"); + amax(blk) = 0.0f; + amax(blk) = max(amax(blk), abs(in[0](r, blk))); + scale(blk) = amax(blk) / (float)qmax; + codes(kk, blk) = cast(round(in[0](kk, blk) / select(scale(blk) == 0.0f, 1.0f, scale(blk)))); + return {codes, scale}; + } + + std::vector decode(const std::vector &encoded) const { + Var kk("kk"), blk("blk"); + Func out("absmax_decoded"); + out(kk, blk) = cast(encoded[0](kk, blk)) * encoded[1](blk); + return {out}; + } + + ApproximationSignature signature() const { + return {{{"block", Float(32), 2}}, + {{"codes", Int(8), 2, ApproximationRange(-qmax, qmax)}, {"scale", Float(32), 1}}}; + } + + Func error_bound(const std::vector &, const std::vector &encoded) const { + Var kk("kk"), blk("blk"); + Func bound("absmax_error_bound"); + bound(kk, blk) = abs(cast(encoded[1](blk))) * Expr(0.5 + qmax / (double)(1 << 21)); + return bound; + } +}; + +Approximation make_quantizer(int block, int qmax = 127) { + return AbsMaxQuantizer{block, qmax}; +} + +// Signed q4 codes [-8, 7] to stored nibbles [0, 15]: exact within that range. +Approximation make_offset() { + return Pointwise{"offset", + [](Expr x) { return cast(x + 8); }, + [](Expr x) { return cast(x - 8); }} + .with_types(Int(8), UInt(8)) + .with_ranges(ApproximationRange(-8, 7), ApproximationRange(0, 15)) + .with_lossless(); +} + +// An int8 to its low nibble (uint8) and high nibble (int8): x = high * 16 + low. +struct SplitNibbles { + std::vector encode(const std::vector &in) const { + Var x("x"); + Func low("split_low"), high("split_high"); + low(x) = cast(in[0](x) & 15); + high(x) = cast(in[0](x) >> 4); + return {low, high}; + } + std::vector decode(const std::vector &in) const { + Var x("x"); + Func out("split_joined"); + out(x) = cast(in[1](x) * 16 + in[0](x)); + return {out}; + } + ApproximationSignature signature() const { + return {{{"codes", Int(8), 1}}, + {{"low", UInt(8), 1, ApproximationRange(0, 15)}, + {"high", Int(8), 1, ApproximationRange(-8, 7)}}}; + } + bool lossless() const { + return true; + } +}; + +int test_round_trip_report() { + // (within, block) input: 4 blocks of 32. + Approximation q = make_quantizer(32); + RoundTripReport r = verify_round_trip(q, Distribution::normal(0, 1), Float(32), {32, 4}, 42); + std::cout << r; + CHECK(r.count == 128); + CHECK(r.seed == 42); + CHECK(r.bound_declared); + CHECK(r.bound_violations == 0); + CHECK(r.max_abs_error > 0 && r.max_abs_error < 0.05); + CHECK(r.rmse > 0 && r.rmse <= r.max_abs_error); + CHECK(!r.distribution.empty()); + + // Adversarial inputs still respect the declared bound. + for (const Distribution &d : {Distribution::zeros(), + Distribution::outliers(Distribution::normal(0, 1), 32, 1000), + Distribution::blockwise_constant(32, Distribution::uniform(-4, 4))}) { + RoundTripReport rr = verify_round_trip(q, d, Float(32), {32, 4}, 5); + CHECK(rr.bound_violations == 0); + } + + // Same seed, same report. + RoundTripReport r2 = verify_round_trip(q, Distribution::normal(0, 1), Float(32), {32, 4}, 42); + CHECK(r2.max_abs_error == r.max_abs_error && r2.rmse == r.rmse); + return 0; +} + +int test_lossless() { + Approximation reshape = BlockReshape{8}; + CHECK(check_property(reshape, lossless(), Distribution::normal(0, 1), Float(32), {64}, 4, 1).passed); + CHECK(check_property(reshape, lossless(), Distribution::uniform_int(-100, 100), Int(16), {64}, 4, 1).passed); + return 0; +} + +// The bit-packing unit at the heart of the motivating example: 4-bit fields +// are exact only in [0, 15]. +int test_precondition_conditioning() { + Approximation pack = PlanarFieldPack{4, 8}; + CHECK(pack.signature().inputs[0].range.has_value()); + + // (a) Inputs generated within the declared precondition. + PropertyResult ok = check_property(pack, lossless(), {InputSpec{UInt(8), {16, 4}, Distribution::uniform_int(0, 15)}}, 4, 3); + std::cout << ok; + CHECK(ok.passed); + + // With no generator given, inputs come from the declared port (type and + // range): here the whole valid range of int8 codes. + Approximation offset = make_offset(); + CHECK(check_property(offset, lossless(), {16, 4}, 4, 3).passed); + + // Outside it: by default the precondition failure is reported ... + std::vector wide = {InputSpec{UInt(8), {16, 4}, Distribution::uniform_int(0, 255)}}; + PropertyResult pre = check_property(pack, lossless(), wide, 4, 3); + CHECK(!pre.passed && pre.precondition_violated); + + // ... and with checking off, the property itself fails (without aborting). + PropertyOptions opts; + opts.check_preconditions = false; + PropertyResult bad = check_property(pack, lossless(), wide, 4, 3, opts); + std::cout << bad; + CHECK(!bad.passed && !bad.precondition_violated); + CHECK(!bad.failing_coord.empty()); + + // The failing seed reproduces the failure in one trial. + PropertyResult again = check_property(pack, lossless(), wide, 1, bad.failing_seed, opts); + CHECK(!again.passed && again.trials_run == 1); + CHECK(again.failing_coord == bad.failing_coord); + CHECK(again.message == bad.message); + return 0; +} + +// Quantize (4-bit, symmetric) -> shift codes to [0, 15] -> pack in 4 bits. +// The packing is exact only because the quantizer guarantees its codes. +int test_stage_targeting() { + Approximation pack = PlanarFieldPack{4, 8}; + Approximation offset = make_offset(); + Approximation quant = make_quantizer(16, 7); + Approximation scheme = Compose{BlockReshape{16}, quant, Parallel{{"codes", Compose{offset, pack}}}}; + + CHECK(check_ranges(scheme).empty()); + std::cout << scheme.describe(); + + // Packing alone, on arbitrary bytes, is not lossless... + // ...but the values that actually reach it here are in [0, 15]. + Distribution normal = Distribution::normal(0, 1); + PropertyResult at_pack = check_property(scheme, lossless().at(pack), normal, Float(32), {64}, 8, 11); + std::cout << at_pack; + CHECK(at_pack.passed); + CHECK(at_pack.stage_label == pack.label()); + + PropertyResult at_offset = check_property(scheme, lossless().at(offset), normal, Float(32), {64}, 8, 11); + CHECK(at_offset.passed); + + CHECK(check_property(scheme, outputs_within_declared_ranges(), normal, Float(32), {64}, 8, 11).passed); + CHECK(check_property(scheme, outputs_within_declared_ranges().at(quant), normal, Float(32), {64}, 8, 11).passed); + return 0; +} + +int test_idempotent_requantize() { + Approximation q = make_quantizer(32); + PropertyResult r = check_property(q, idempotent_requantize(), Distribution::normal(0, 1), Float(32), {32, 4}, 6, 9); + std::cout << r; + CHECK(r.passed); + CHECK(check_property(q, within_declared_bound(), Distribution::normal(0, 3), Float(32), {32, 4}, 6, 9).passed); + CHECK(check_property(q, outputs_within_declared_ranges(), Distribution::normal(0, 3), Float(32), {32, 4}, 6, 9).passed); + CHECK(check_property(q, zero_preserving(), Distribution::zeros(), Float(32), {32, 4}, 2, 1).passed); + CHECK(check_property(q, sign_preserving(), Distribution::normal(0, 1), Float(32), {32, 4}, 4, 1).passed); + CHECK(!check_property(q, lossless(), Distribution::normal(0, 1), Float(32), {32, 4}, 2, 1).passed); + CHECK(!check_property(q, bounded_error(1e-9), Distribution::normal(0, 1), Float(32), {32, 4}, 2, 1).passed); + return 0; +} + +int test_range_diagnostics() { + // Codes [-8, 8] are one too many for the offset's precondition [-8, 7]. + Approximation pack = PlanarFieldPack{4, 8}; + Approximation offset = make_offset(); + Approximation quant = make_quantizer(16, 8); + Approximation bad = Compose{BlockReshape{16}, quant, Parallel{{"codes", Compose{offset, pack}}}}; + std::vector issues = check_ranges(bad); + for (const std::string &s : issues) { + std::cout << "issue: " << s << "\n"; + } + CHECK(!issues.empty()); + std::string text = bad.describe(); + std::cout << text; + CHECK(text.find("not within") != std::string::npos); + return 0; +} + +int test_declared_bound_is_exact() { + // The declared bound holds on hostile input. + for (int qmax : {7, 127}) { + Approximation q = make_quantizer(16, qmax); + for (const Distribution &d : {Distribution::normal(0, 1), Distribution::uniform(-1, 1), + Distribution::outliers(Distribution::uniform(-1, 1), 16, 50)}) { + CHECK(check_property(q, within_declared_bound(), d, Float(32), {16, 4}, 6, 21).passed); + CHECK(check_property(q, outputs_within_declared_ranges(), d, Float(32), {16, 4}, 6, 21).passed); + } + } + return 0; +} + +// Copies a (possibly struct-typed) Func unchanged. +struct Copy { + Func encode(const Func &f) const { + Func g("copy_encode"); + g(_) = f(_); + return g; + } + Func decode(const Func &f) const { + Func g("copy_decode"); + g(_) = f(_); + return g; + } +}; + +// A Tuple-valued encoded Func: (x, x + 1). +struct TuplePair { + Func encode(const Func &f) const { + Func t("tuple_pair"); + t(_) = Tuple(f(_), f(_) + cast(1)); + return t; + } + Func decode(const Func &f) const { + Func g("tuple_first"); + g(_) = f(_)[0]; + return g; + } + ApproximationSignature signature() const { + return {{{"value", UInt(8), 1}}, {{"pair", std::nullopt, 1, ApproximationRange(0, 255)}}}; + } + bool lossless() const { + return true; + } +}; + +bool mentions_unreadable(const PropertyResult &r) { + return !r.passed && r.message.find("cannot be read back") != std::string::npos; +} + +// Encoded Funcs that are struct-typed or Tuple-valued cannot be read back: +// properties that need their values fail with a message; the others work. +int test_unreadable_encoded() { + Type record = Type::Struct({{"low", UInt(8)}, {"high", Int(8)}}); + Approximation split = SplitNibbles{}; + Approximation layout = StructLayout(record, {"low", "high"}); + Approximation copy = Copy{}; + Approximation scheme = Compose{split, layout}; + + const std::vector extents = {32}; + const std::vector codes = {InputSpec{Int(8), extents, Distribution::uniform_int(-16, 15)}}; + CHECK(scheme.lossless()); + PropertyResult lossless_result = check_property(scheme, lossless(), codes, 3, 5); + std::cout << lossless_result; + CHECK(lossless_result.passed); + CHECK(check_property(scheme, bounded_error(0), codes, 2, 5).passed); + CHECK(check_property(scheme, within_declared_bound(), codes, 2, 5).passed); + CHECK(check_property(scheme, zero_preserving(), codes, 1, 5).passed); + CHECK(check_property(scheme, sign_preserving(), codes, 2, 5).passed); + + // The record port declares no range, so there is nothing to check. + CHECK(check_property(scheme, outputs_within_declared_ranges(), codes, 2, 5).passed); + + PropertyResult idem = check_property(scheme, idempotent_requantize(), codes, 2, 5); + std::cout << idem; + CHECK(mentions_unreadable(idem)); + CHECK(idem.message.find("'record'") != std::string::npos && idem.message.find("struct-typed") != std::string::npos); + + // A stage whose input is struct-typed cannot be targeted, but others can. + Approximation with_copy = Compose{split, layout, copy}; + CHECK(check_property(with_copy, lossless(), codes, 2, 5).passed); + PropertyResult at_copy = check_property(with_copy, lossless().at(copy), codes, 2, 5); + std::cout << at_copy; + CHECK(mentions_unreadable(at_copy)); + CHECK(check_property(with_copy, lossless().at(split), codes, 2, 5).passed); + CHECK(check_property(with_copy, lossless().at(layout), codes, 2, 5).passed); + + // The layout on its own, over two inputs of different types. + Approximation two = StructLayout(record, {"low", "high"}); + std::vector specs = {InputSpec{UInt(8), extents, Distribution::uniform_int(0, 255)}, + InputSpec{Int(8), extents, Distribution::uniform_int(-128, 127)}}; + CHECK(check_property(two, lossless(), specs, 2, 5).passed); + RoundTripReport report = verify_round_trip(two, {generate(specs[0].dist, UInt(8), extents, 1), + generate(specs[1].dist, Int(8), extents, 2)}); + CHECK(report.max_abs_error == 0); + + // Tuple-valued. + Approximation tuple = TuplePair{}; + CHECK(check_property(tuple, lossless(), extents, 2, 5).passed); + CHECK(mentions_unreadable(check_property(tuple, idempotent_requantize(), extents, 2, 5))); + PropertyResult ranges = check_property(tuple, outputs_within_declared_ranges(), extents, 2, 5); + CHECK(mentions_unreadable(ranges) && ranges.message.find("Tuple-valued") != std::string::npos); + return 0; +} + +// SplitNibbles' declared output ranges hold for every code. +int test_split_ranges() { + Approximation split = SplitNibbles{}; + CHECK(check_property(split, outputs_within_declared_ranges(), {InputSpec{Int(8), {256}, Distribution::uniform_int(-128, 127)}}, 4, 3).passed); + CHECK(check_property(split, lossless(), {InputSpec{Int(8), {256}, Distribution::uniform_int(-128, 127)}}, 4, 3).passed); + CHECK(check_property(split, idempotent_requantize(), {InputSpec{Int(8), {256}, Distribution::uniform_int(-128, 127)}}, 2, 3).passed); + return 0; +} + +} // namespace + +int main(int argc, char **argv) { + int (*tests[])() = {test_generators, test_round_trip_report, test_lossless, test_precondition_conditioning, + test_stage_targeting, test_idempotent_requantize, test_range_diagnostics, + test_declared_bound_is_exact, test_unreadable_encoded, test_split_ranges}; + for (auto t : tests) { + if (t()) { + return 1; + } + } + printf("Success!\n"); + return 0; +} diff --git a/test/correctness/approximation_trace.cpp b/test/correctness/approximation_trace.cpp new file mode 100644 index 000000000000..acdf2c911d0d --- /dev/null +++ b/test/correctness/approximation_trace.cpp @@ -0,0 +1,206 @@ +#include "Halide.h" +#include +#include + +using namespace Halide; + +namespace { + +struct Named { + Func encode(const Func &f) const { + return f; + } + Func decode(const Func &f) const { + return f; + } + std::string name() const { + return "MyName"; + } +}; + +// A user unit that runs two handles inside its own encode()/decode(). +struct Pair { + Approximation first, second; + + std::vector encode(const std::vector &inputs) const { + return second.encode(first.encode(inputs).encoded).encoded; + } + std::vector decode(const std::vector &encoded) const { + return first.decode(second.decode(encoded).decoded).decoded; + } + std::string name() const { + return "Pair"; + } +}; + +// Drop "$N" uniquifier suffixes so the expected text is stable. +std::string normalize(const std::string &s) { + std::string out; + for (size_t i = 0; i < s.size(); i++) { + if (s[i] == '$') { + while (i + 1 < s.size() && isdigit((unsigned char)s[i + 1])) { + i++; + } + } else { + out += s[i]; + } + } + return out; +} + +#define CHECK(cond) \ + do { \ + if (!(cond)) { \ + printf("%s:%d: check failed: %s\n", __FILE__, __LINE__, #cond); \ + return 1; \ + } \ + } while (0) + +std::vector names(const std::vector &fs) { + std::vector result; + for (const Func &f : fs) { + result.push_back(f.name()); + } + return result; +} + +} // namespace + +int main() { + // Labels. + { + Approximation comp = Identity{}; + CHECK(comp.label() == "Identity"); + CHECK(Approximation(BlockReshape{32}).label() == "BlockReshape"); + CHECK(Approximation(Compose{Identity{}, Identity{}}).label() == "Compose"); + CHECK(Approximation(Parallel{Identity{}, Identity{}}).label() == "Parallel"); + CHECK(Approximation(LittleEndianScalarPack{}).label() == "LittleEndianScalarPack"); + CHECK(Approximation().label().empty()); + + Approximation named = Named{}; + CHECK(named.label() == "MyName"); + + Approximation copy = named; + Approximation renamed = named.labelled("custom"); + CHECK(renamed.same_as(named)); + CHECK(named.label() == "custom"); // shared state + CHECK(copy.label() == "custom"); + CHECK(renamed.label() == "custom"); + + Approximation ctor(Named{}, "explicit"); + CHECK(ctor.label() == "explicit"); + CHECK(!ctor.same_as(named)); + } + + // Tree shape, printing, stage_ports. + Func src("src"); + Var x("x"); + src(x) = x; + Func consumer("consumer"); + consumer(x) = src(x) + 1; + + Approximation neg = Pointwise{"neg", [](Expr e) { return -e; }, [](Expr e) { return -e; }}; + Approximation inc = Pointwise{"inc", [](Expr e) { return e + 1; }, [](Expr e) { return e - 1; }}; + Approximation dbl(Pointwise{"dbl", [](Expr e) { return e * 2; }, [](Expr e) { return e / 2; }}, "Double"); + Approximation pair = Pair{inc, dbl}; + Approximation scheme = Compose{neg, Parallel{std::vector{pair}}}; + + ApproximationResult r = src.approximate_by(scheme, {consumer}); + + const ApproximationTraceNode &et = r.encode_trace; + CHECK(et.stage.same_as(scheme)); + CHECK(et.label == "Compose"); + // Compose encodes front-to-back: neg first, then the Parallel. + CHECK(et.children.size() == 2); + CHECK(et.children[0].stage.same_as(neg)); + CHECK(et.children[0].children.empty()); + CHECK(et.children[1].label == "Parallel"); + CHECK(et.children[1].children.size() == 1); + const ApproximationTraceNode &ep = et.children[1].children[0]; + CHECK(ep.stage.same_as(pair)); + CHECK(ep.children.size() == 2); + CHECK(ep.children[0].stage.same_as(inc)); + CHECK(ep.children[1].stage.same_as(dbl)); + CHECK(ep.children[1].label == "Double"); + + // Decode runs back-to-front, and Pair undoes its stages in reverse. + const ApproximationTraceNode &dt = r.decode_trace; + CHECK(dt.stage.same_as(scheme)); + CHECK(dt.children.size() == 2); + CHECK(dt.children[0].label == "Parallel"); + CHECK(dt.children[1].stage.same_as(neg)); + const ApproximationTraceNode &dp = dt.children[0].children[0]; + CHECK(dp.stage.same_as(pair)); + CHECK(dp.children.size() == 2); + CHECK(dp.children[0].stage.same_as(dbl)); + CHECK(dp.children[1].stage.same_as(inc)); + + // The flat view is the post-order flattening of the tree, and the + // lookups still work. + CHECK(r.encoded_stage_outputs.size() == 6); + CHECK(r.encoded_stage_outputs[0].stage.same_as(neg)); + CHECK(r.encoded_stage_outputs[5].stage.same_as(scheme)); + CHECK(r.encoded_by(inc).name() == "inc_encode"); + CHECK(r.decoded_by(dbl).name() == "dbl_decode"); + CHECK(r.decoded_by(neg).name() == r.replacement.name()); + + std::ostringstream trace; + trace << et; + const char *expected_trace = + "Compose -> 0=dbl_encode\n" + " intermediates: neg_encode, inc_encode\n" + " Pointwise -> 0=neg_encode\n" + " Parallel -> 0=dbl_encode\n" + " intermediates: inc_encode\n" + " Pair -> 0=dbl_encode\n" + " intermediates: inc_encode\n" + " Pointwise -> 0=inc_encode\n" + " Double -> 0=dbl_encode\n"; + if (normalize(trace.str()) != expected_trace) { + printf("Unexpected trace:\n%s", trace.str().c_str()); + return 1; + } + + std::ostringstream full; + full << r; + const char *expected_full = + "encode:\n" + " Compose -> 0=dbl_encode\n" + " intermediates: neg_encode, inc_encode\n" + " Pointwise -> 0=neg_encode\n" + " Parallel -> 0=dbl_encode\n" + " intermediates: inc_encode\n" + " Pair -> 0=dbl_encode\n" + " intermediates: inc_encode\n" + " Pointwise -> 0=inc_encode\n" + " Double -> 0=dbl_encode\n" + "decode:\n" + " Compose -> 0=neg_decode\n" + " intermediates: dbl_decode, inc_decode\n" + " Parallel -> 0=inc_decode\n" + " intermediates: dbl_decode\n" + " Pair -> 0=inc_decode\n" + " intermediates: dbl_decode\n" + " Double -> 0=dbl_decode\n" + " Pointwise -> 0=inc_decode\n" + " Pointwise -> 0=neg_decode\n"; + if (normalize(full.str()) != expected_full) { + printf("Unexpected result printout:\n%s", full.str().c_str()); + return 1; + } + + // Trace order, encode side first; the replacement (neg_decode) is excluded. + const std::vector expected_ports = { + "neg_encode", "inc_encode", "dbl_encode", "dbl_decode", "inc_decode"}; + std::vector got = names(r.stage_ports()); + for (std::string &n : got) { + n = normalize(n); + } + CHECK(got == expected_ports); + CHECK(r.is_stage_port(r.encoded_by(inc))); + CHECK(!r.is_stage_port(r.replacement)); + CHECK(!r.is_stage_port(consumer)); + + printf("Success!\n"); + return 0; +} diff --git a/test/correctness/approximation_unit_forms.cpp b/test/correctness/approximation_unit_forms.cpp new file mode 100644 index 000000000000..20591708a949 --- /dev/null +++ b/test/correctness/approximation_unit_forms.cpp @@ -0,0 +1,254 @@ +#include "Halide.h" +#include + +using namespace Halide; + +namespace { + +// Single-form encode/decode. +struct Negate { + Func encode(const Func &f) const { + Func g("negated"); + g(_) = -f(_); + return g; + } + Func decode(const Func &f) const { + return encode(f); + } +}; + +// Multi-form: two Funcs out, back into one. +struct SplitParity { + std::vector encode(const std::vector &in) const { + Var x("x"); + Func even("even"), odd("odd"); + even(x) = in[0](2 * x); + odd(x) = in[0](2 * x + 1); + return {even, odd}; + } + std::vector decode(const std::vector &in) const { + Var x("x"); + Func out("merged"); + out(x) = select(x % 2 == 0, in[0](x / 2), in[1](x / 2)); + return {out}; + } +}; + +// Mixed: multi encode + single decode. +struct AddOne { + std::vector encode(const std::vector &in) const { + Var x("x"); + Func g("plus_one"); + g(x) = in[0](x) + 1; + return {g}; + } + Func decode(const Func &f) const { + Var x("x"); + Func g("minus_one"); + g(x) = f(x) - 1; + return g; + } +}; + +// Both a Func and a vector overload: the vector one must win. +struct BothOverloads { + std::vector encode(const std::vector &in) const { + return in; + } + Func encode(const Func &f) const { + return f; + } + std::vector decode(const std::vector &in) const { + return in; + } + Func decode(const Func &f) const { + return f; + } +}; + +struct EncodeOnly { + Func encode(const Func &f) const { + return f; + } +}; + +struct BadReturn { + int encode(const Func &) const { + return 0; + } + Func decode(const Func &f) const { + return f; + } +}; + +struct NonConstEncode { + Func encode(const Func &f) { + return f; + } + Func decode(const Func &f) const { + return f; + } +}; + +static_assert(std::is_convertible_v); +static_assert(std::is_convertible_v); +static_assert(std::is_convertible_v); +static_assert(std::is_convertible_v); +static_assert(std::is_convertible_v); +static_assert(std::is_convertible_v); +// The unit-side full form is gone: results are not a valid unit return type. +struct FullForm { + EncodeResult encode(const std::vector &in) const { + return {in, {}, {}}; + } + DecodeResult decode(const std::vector &in) const { + return {in, {}, {}}; + } +}; + +// A unit that internally calls two member handles. +struct Pair { + Approximation a, b; + std::vector encode(const std::vector &in) const { + return b.encode(a.encode(in).encoded).encoded; + } + std::vector decode(const std::vector &in) const { + return a.decode(b.decode(in).decoded).decoded; + } +}; + +static_assert(!std::is_convertible_v); +static_assert(!std::is_convertible_v); +static_assert(!std::is_convertible_v); +static_assert(!std::is_convertible_v); +static_assert(!std::is_convertible_v); + +template +int check_round_trip(const char *what, const Approximation &a, F expected, int n = 16) { + Func f("f"); + Var x("x"); + f(x) = x * 3 + 1; + Func g("g"); + g(x) = f(x) * 2; + ApproximationResult r = f.approximate_by(a, {g}); + for (Func h : r.intermediates) { + h.compute_root(); + } + r.replacement.compute_root(); + Buffer out = g.realize({n}); + for (int i = 0; i < n; i++) { + int want = expected(i) * 2; + if (out(i) != want) { + printf("%s: g(%d) = %d, expected %d\n", what, i, out(i), want); + return 1; + } + } + return 0; +} + +} // namespace + +int main() { + auto identity = [](int i) { return i * 3 + 1; }; + + Approximation neg = Negate{}; + if (check_round_trip("single", neg, identity)) return 1; + + Approximation parity = SplitParity{}; + if (check_round_trip("multi", parity, identity)) return 1; + + Approximation add = AddOne{}; + if (check_round_trip("mixed", add, identity)) return 1; + + if (check_round_trip("both", BothOverloads{}, identity)) return 1; + + Approximation pw = Pointwise{"scale", + [](Expr v) { return v * 2; }, + [](Expr v) { return v / 2; }}; + if (check_round_trip("pointwise", pw, identity)) return 1; + + // Tuple-valued Pointwise: swap the halves of a pair. + { + Var x("x"); + Func t("t"); + t(x) = Tuple(x, x * 10); + Approximation swap = Pointwise{"swap", + [](const std::vector &v) { return std::vector{v[1], v[0]}; }, + [](const std::vector &v) { return std::vector{v[1], v[0]}; }}; + EncodeResult e = swap.encode({t}); + if (e.encoded[0].name().rfind("swap_encode", 0) != 0) { + printf("unexpected pointwise name %s\n", e.encoded[0].name().c_str()); + return 1; + } + DecodeResult d = swap.decode(e.encoded); + Realization re = d.decoded[0].realize({4}); + Buffer a = re[0], b = re[1]; + for (int i = 0; i < 4; i++) { + if (a(i) != i || b(i) != i * 10) { + printf("tuple pointwise mismatch at %d\n", i); + return 1; + } + } + } + + // Stage lookup works for simple-form units, including when composed. + { + Approximation h_neg = Negate{}, h_parity = SplitParity{}, h_add = AddOne{}; + Compose scheme{h_add, h_neg, h_parity}; + Func f("src"); + Var x("x"); + f(x) = x; + Func c("consumer"); + c(x) = f(x); + ApproximationResult r = f.approximate_by(scheme, {c}); + auto starts_with = [](const Func &fn, const char *prefix) { + return fn.defined() && fn.name().rfind(prefix, 0) == 0; + }; + if (!starts_with(r.decoded_by(h_neg), "negated") || !starts_with(r.encoded_by(h_add), "plus_one") || + !starts_with(r.decoded_by(h_add), "minus_one") || !r.decoded_by(h_parity).defined()) { + printf("stage lookup returned unexpected Funcs\n"); + return 1; + } + } + + // A hand-written unit that calls other handles is traced automatically. + { + Approximation inner_neg = Negate{}, inner_add = AddOne{}; + Approximation pair = Pair{inner_neg, inner_add}; + Func f("pair_src"); + Var x("x"); + f(x) = x; + EncodeResult e = pair.encode({f}); + DecodeResult d = pair.decode(e.encoded); + auto names = [](const std::vector &so) { + std::string s; + for (const ApproximationStageOutputs &o : so) { + s += o.ports[0].name() + ","; + } + return s; + }; + if (e.stage_outputs.size() != 3 || !e.stage_outputs[0].stage.same_as(inner_neg) || + !e.stage_outputs[1].stage.same_as(inner_add) || !e.stage_outputs[2].stage.same_as(pair) || + d.stage_outputs.size() != 3 || !d.stage_outputs[0].stage.same_as(inner_add) || + !d.stage_outputs[1].stage.same_as(inner_neg) || !d.stage_outputs[2].stage.same_as(pair)) { + printf("nested unit trace wrong: enc [%s] dec [%s]\n", names(e.stage_outputs).c_str(), + names(d.stage_outputs).c_str()); + return 1; + } + // The pair's inter-stage Func is discovered as an intermediate. + if (e.intermediates.size() != 1 || e.intermediates[0].name().rfind("negated", 0) != 0) { + printf("nested unit intermediates wrong\n"); + return 1; + } + Buffer out = d.decoded[0].realize({8}); + for (int i = 0; i < 8; i++) { + if (out(i) != i) { + printf("nested unit round trip wrong at %d\n", i); + return 1; + } + } + } + + printf("Success!\n"); + return 0; +} diff --git a/test/correctness/hoist_invariants.cpp b/test/correctness/hoist_invariants.cpp index 7f61562ca580..d4457f9a190e 100644 --- a/test/correctness/hoist_invariants.cpp +++ b/test/correctness/hoist_invariants.cpp @@ -588,6 +588,93 @@ int hoist_invariants_nothing_to_hoist_rejected_test() { return 0; } +// distribute() multiplies a product of sums out so that hoist_invariants() can +// see terms that were not written as terms. This is the affine-quantized dot +// product: sum_k (d*q_k + m) * (e*p_k), whose expansion +// d*e*sum_k(q_k*p_k) + m*e*sum_k(p_k) has two integer-bodied accumulators where +// the unexpanded form has one float one. +int hoist_invariants_distribute_test() { + const int K = 32; + ImageParam Q{Int(8), 1, "Q"}; + ImageParam P{Int(8), 1, "P"}; + ImageParam D{Float(32), 1, "D"}; + ImageParam M{Float(32), 1, "M"}; + ImageParam E{Float(32), 1, "E"}; + + Var i{"i"}; + RDom r(0, K, "r"); + + Func Acc{"Acc"}; + Acc(i) = 0.0f; + Acc(i) += (cast(Q(r)) * D(i) + M(i)) * (cast(P(r)) * E(i)); + + std::vector Acc_intm = Acc.update().distribute().hoist_invariants(); + internal_assert(Acc_intm.size() == 2) + << "distribute: expected the multiplied-out increment to yield two " + << "accumulators, got " << Acc_intm.size() << "\n"; + + // Both bodies are integer sums of int8 products, so both retype -- and each + // accumulator is its own Func, so each may take its own target type. + Func qp = Acc_intm[0].change_type(Int(32)); + Func p_sum = Acc_intm[1].change_type(Int(16)); + internal_assert(qp.types()[0] == Int(32) && p_sum.types()[0] == Int(16)) + << "distribute: retyping the accumulators separately gave " + << qp.types()[0] << " and " << p_sum.types()[0] << "\n"; + qp.compute_root(); + p_sum.compute_root(); + + Buffer q_buf(K), p_buf(K); + Buffer d_buf(1), m_buf(1), e_buf(1); + int64_t sum_qp = 0, sum_p = 0; + for (int k = 0; k < K; k++) { + q_buf(k) = (int8_t)(k % 15); + p_buf(k) = (int8_t)((k * 5) % 127 - 63); + sum_qp += (int64_t)q_buf(k) * p_buf(k); + sum_p += p_buf(k); + } + d_buf(0) = 0.25f; + m_buf(0) = -0.5f; + e_buf(0) = 2.0f; + Q.set(q_buf); + P.set(p_buf); + D.set(d_buf); + M.set(m_buf); + E.set(e_buf); + + Buffer result = Acc.realize({1}); + const float expected = 0.25f * 2.0f * (float)sum_qp + -0.5f * 2.0f * (float)sum_p; + internal_assert(std::abs(result(0) - expected) < 1e-3f) + << "distribute: got " << result(0) << ", expected " << expected << "\n"; + + return 0; +} + +// distribute() is a schedule decision, so it says so when there is nothing to +// multiply out rather than quietly leaving the reduction alone. +int distribute_nothing_to_do_rejected_test() { + if (!Halide::exceptions_enabled()) { + return 0; + } + ImageParam A{Float(32), 1, "A"}; + Var i{"i"}; + RDom r(0, 8, "r"); + + Func f{"f"}; + f(i) = 0.0f; + f(i) += A(i) * cast(r); + + try { + f.update().distribute(); + } catch (const Halide::CompileError &e) { + const std::string msg = e.what(); + internal_assert(msg.find("no product over a sum") != std::string::npos) + << "distribute() rejected the update for the wrong reason: " << msg << "\n"; + return 0; + } + internal_assert(false) << "distribute() accepted an update with nothing to multiply out\n"; + return 0; +} + } // namespace int main(int argc, char **argv) { @@ -610,6 +697,8 @@ int main(int argc, char **argv) { {"hoist_invariants test (after rfactor)", hoist_invariants_after_rfactor_test}, {"hoist_invariants test (invalid law rejected)", hoist_invariants_invalid_law_rejected_test}, {"hoist_invariants test (nothing to hoist rejected)", hoist_invariants_nothing_to_hoist_rejected_test}, + {"distribute test (affine dot product)", hoist_invariants_distribute_test}, + {"distribute test (nothing to distribute rejected)", distribute_nothing_to_do_rejected_test}, }; using Sharder = Halide::Internal::Test::Sharder; diff --git a/test/correctness/parallel_fork.cpp b/test/correctness/parallel_fork.cpp index 5183370072a4..fe79125cb10c 100644 --- a/test/correctness/parallel_fork.cpp +++ b/test/correctness/parallel_fork.cpp @@ -21,7 +21,7 @@ namespace halide_externs { HalideExtern_1(int, five_ms, int); } -enum Schedule { +enum class Schedule { Serial, Parallel, AsyncRoot, @@ -40,20 +40,20 @@ Func make(Schedule schedule) { both.compute_root().bound(z, 0, 2); switch (schedule) { - case Serial: + case Schedule::Serial: f.compute_root(); g.compute_root(); break; - case Parallel: + case Schedule::Parallel: both.parallel(z); f.compute_at(both, z); g.compute_at(both, z); break; - case AsyncRoot: + case Schedule::AsyncRoot: f.compute_root().async(); g.compute_root().async(); break; - case AsyncComputeAt: + case Schedule::AsyncComputeAt: both.parallel(z); f.compute_at(both, z).async(); g.compute_at(both, z).async(); @@ -75,7 +75,7 @@ int main(int argc, char **argv) { double time; call_count = 0; - both = make(Serial); + both = make(Schedule::Serial); im = both.realize({10, 10, 2}); count = call_count; time = benchmark([&]() { @@ -85,7 +85,7 @@ int main(int argc, char **argv) { fflush(stdout); call_count = 0; - both = make(Parallel); + both = make(Schedule::Parallel); im = both.realize({10, 10, 2}); count = call_count; time = benchmark([&]() { @@ -94,7 +94,7 @@ int main(int argc, char **argv) { printf("Parallel time %f for %d calls.\n", time, count); fflush(stdout); - both = make(AsyncRoot); + both = make(Schedule::AsyncRoot); call_count = 0; im = both.realize({10, 10, 2}); count = call_count; @@ -104,7 +104,7 @@ int main(int argc, char **argv) { printf("Async root time %f for %d calls.\n", time, count); fflush(stdout); - both = make(AsyncComputeAt); + both = make(Schedule::AsyncComputeAt); call_count = 0; im = both.realize({10, 10, 2}); count = call_count; diff --git a/test/correctness/sever.cpp b/test/correctness/sever.cpp index 94b95ac278e3..ea9fd804701e 100644 --- a/test/correctness/sever.cpp +++ b/test/correctness/sever.cpp @@ -207,6 +207,96 @@ int named_bindings_test() { return 0; } +// The same quantizer expressed as an Approximation, used to check that +// sever() composes with approximate_by()'s rewritten call graph. +struct ApproxSymmetricQuantize { + explicit ApproxSymmetricQuantize(int k) + : k_(k) { + } + + std::vector encode(const std::vector &inputs) const { + Func v = inputs[0]; + Var k("k"); + RDom r(0, k_, "r"); + + Func amax("amax"); + amax() = 0.0f; + amax() = max(amax(), abs(v(r))); + + Func d("scale"); + d() = amax() / 127.0f; + + Func q("q"); + Expr id = select(d() != 0.0f, 1.0f / d(), 0.0f); + q(k) = cast(clamp(round(v(k) * id), -127, 127)); + + return {q, d}; + } + + std::vector decode(const std::vector &encoded) const { + Func q = encoded[0], d = encoded[1]; + Var k("k"); + Func dequantized("dequantized"); + dequantized(k) = cast(q(k)) * d(); + return {dequantized}; + } + +private: + int k_; +}; + +// Sever a quantized vector's encode() outputs from a consumer built via +// approximate_by(), and check the final result still matches the plain-C++ +// reference round trip. +int approximate_by_offline_test() { + const int K = 64; + Var k("k"); + + Func Vec("Vec"); + Vec(k) = cos(cast(k) * 0.05f) * 3.0f; + + ApproxSymmetricQuantize quantize(K); + + Func Result("Result"); + Result(k) = Vec(k) * 2.0f; + + ApproximationResult result = Vec.approximate_by(quantize, {Result}); + Result.eager_inline({result.replacement}); + + // result.intermediates is [q, d, amax]: encode()'s two signature-contract + // outputs, then the scheduling-only intermediate the framework discovered. q and d are the actual + // Funcs Result's call graph depends on (approximate_by() calls encode() + // internally; a separately-called quantize.encode({Vec}) here would + // build an unrelated, unconnected copy of the same graph shape). + std::vector encoded = {result.intermediates[0], result.intermediates[1]}; + for (size_t i = 2; i < result.intermediates.size(); i++) { + result.intermediates[i].compute_root(); + } + + SeverResult split = Pipeline({Result}).sever(encoded); + + Buffer q_buf(K); + Buffer scale_buf = Buffer::make_scalar(); + split.offline.realize({q_buf, scale_buf}); + split.online_inputs[0].set(q_buf); + split.online_inputs[1].set(scale_buf); + + Buffer out = Result.realize({K}); + + std::vector ref_q; + float ref_scale; + reference_symmetric_quantize(K, [](int kk) { return cosf(kk * 0.05f) * 3.0f; }, ref_q, ref_scale); + for (int kk = 0; kk < K; kk++) { + float expected = (ref_q[kk] * ref_scale) * 2.0f; + if (std::fabs(out(kk) - expected) > 1e-3f * std::fabs(expected)) { + printf("approximate_by_offline_test: Result(%d) = %f, expected %f\n", kk, out(kk), expected); + return 1; + } + } + + return 0; +} + } // namespace int main(int argc, char **argv) { @@ -219,6 +309,9 @@ int main(int argc, char **argv) { if (named_bindings_test()) { return 1; } + if (approximate_by_offline_test()) { + return 1; + } printf("Success!\n"); return 0; diff --git a/tools/CMakeLists.txt b/tools/CMakeLists.txt index b0d510dfc7a1..a66fd1f7ce07 100644 --- a/tools/CMakeLists.txt +++ b/tools/CMakeLists.txt @@ -91,6 +91,7 @@ target_sources( INTERFACE FILE_SET HEADERS FILES + halide_approximation_testing.h halide_benchmark.h halide_image.h halide_image_info.h diff --git a/tools/halide_approximation_testing.h b/tools/halide_approximation_testing.h new file mode 100644 index 000000000000..b44c21c186f4 --- /dev/null +++ b/tools/halide_approximation_testing.h @@ -0,0 +1,1523 @@ +#ifndef HALIDE_APPROXIMATION_TESTING_TOOL_H +#define HALIDE_APPROXIMATION_TESTING_TOOL_H + +/** \file + * Property-based testing and round-trip verification for Approximations. + * Header-only, like halide_image_io.h: include it alongside Halide.h in test + * code. Everything lives in Halide::ApproximationTesting, since names like + * Distribution and Property are too generic for namespace Halide. + * + * The pieces: + * + * - Distribution and generate(): seeded, platform-reproducible input + * generators, including adversarial ones (see below). + * - verify_round_trip(): run decode(encode(x)) on concrete inputs and report + * error statistics, checked against the unit's declared error bound if it + * has one. + * - Property and check_property(): named checks (lossless(), bounded_error(), + * idempotent_requantize(), ...) run over several seeded trials. A property + * is *conditioned* on where its inputs come from in two ways: + * (a) explicitly, by choosing the Distribution the inputs are drawn from + * (by default one that respects the root's declared input ranges); + * (b) by stage: `prop.at(stage)` runs the whole approximation on generated + * inputs, takes the values that actually arrive at `stage`'s encode + * inputs, and checks the property for that stage alone on those + * values. A bit-packing stage that is only exact for inputs that fit + * in its field can thus be tested on what an upstream quantizer really + * hands it. + * + * Everything that runs Halide code uses the JIT, with every stage boundary + * compute_root'd and traced so its values can be read back; nothing here + * changes how a pipeline you build yourself is scheduled. + */ + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "Halide.h" + +namespace Halide { +namespace ApproximationTesting { + +// --------------------------------------------------------------------------- +// Random numbers +// --------------------------------------------------------------------------- + +/** One step of the splitmix64 generator: advances `state` and returns the + * next output. (splitmix64(seed = 0) yields 0xe220a8397b1dcdaf first.) */ +inline uint64_t splitmix64(uint64_t &state) { + uint64_t z = (state += 0x9e3779b97f4a7c15ULL); + z = (z ^ (z >> 30)) * 0xbf58476d1ce4e5b9ULL; + z = (z ^ (z >> 27)) * 0x94d049bb133111ebULL; + return z ^ (z >> 31); +} + +/** xoshiro256** seeded through splitmix64. Only integer operations, so a + * given seed produces the same stream on every platform and standard + * library (unlike std::uniform_*_distribution). */ +class Rng { +public: + explicit Rng(uint64_t seed) { + for (uint64_t &word : s_) { + word = splitmix64(seed); + } + } + + uint64_t next() { + const uint64_t result = rotl(s_[1] * 5, 7) * 9; + const uint64_t t = s_[1] << 17; + s_[2] ^= s_[0]; + s_[3] ^= s_[1]; + s_[1] ^= s_[2]; + s_[0] ^= s_[3]; + s_[2] ^= t; + s_[3] = rotl(s_[3], 45); + return result; + } + + /** Uniform in [0, 1), with 53 random bits. */ + double next_double() { + return (double)(next() >> 11) * (1.0 / 9007199254740992.0); + } + +private: + static uint64_t rotl(uint64_t x, int k) { + return (x << k) | (x >> (64 - k)); + } + uint64_t s_[4]; +}; + +/** The seed of the `index`th trial or input derived from `seed`. Index 0 is + * `seed` itself, so re-running with `seed = failing_seed` and one trial + * reproduces a failing trial exactly. */ +inline uint64_t derive_seed(uint64_t seed, uint64_t index) { + if (index == 0) { + return seed; + } + uint64_t state = seed ^ (index * 0xd1342543de82ef95ULL); + splitmix64(state); + return splitmix64(state); +} + +namespace Detail { + +// A product kept out of fused multiply-adds, whose different rounding on +// different compilers would otherwise break bit-reproducibility. +inline double product(double a, double b) { + volatile double p = a * b; + return p; +} + +struct FloatInfo { + double max, min_normal, denorm_min; +}; + +inline FloatInfo float_info(Type t) { + if (t.is_bfloat()) { + return {3.3895313892515355e38, 1.1754943508222875e-38, 9.1835496157991212e-41}; + } + switch (t.bits()) { + case 16: + return {65504.0, 6.103515625e-05, 5.9604644775390625e-08}; + case 32: + return {3.4028234663852886e38, 1.1754943508222875e-38, 1.401298464324817e-45}; + default: + return {std::numeric_limits::max(), std::numeric_limits::min(), + std::numeric_limits::denorm_min()}; + } +} + +inline bool is_floating(Type t) { + return t.is_float() || t.is_bfloat(); +} + +inline bool is_integral(Type t) { + return t.is_int() || t.is_uint(); +} + +// The lowest and highest value of an arithmetic type, as doubles (64-bit +// integer limits round to +-2^63 / 2^64). +inline std::pair type_limits(Type t) { + if (is_floating(t)) { + double m = float_info(t).max; + return {-m, m}; + } + if (t.is_bool() || t.bits() == 1) { + return {0, 1}; + } + if (t.is_uint()) { + return {0, std::ldexp(1.0, t.bits()) - (t.bits() < 53 ? 1 : 0)}; + } + return {-std::ldexp(1.0, t.bits() - 1), std::ldexp(1.0, t.bits() - 1) - (t.bits() < 54 ? 1 : 0)}; +} + +template +T saturate(double v) { + if (!(v == v)) { + return 0; + } + if (v <= (double)std::numeric_limits::lowest()) { + return std::numeric_limits::lowest(); + } + if (v >= (double)std::numeric_limits::max()) { + return std::numeric_limits::max(); + } + return (T)v; +} + +inline void store_value(void *p, Type t, double v) { + if (t.is_float()) { + if (t.bits() == 64) { + *(double *)p = v; + } else if (t.bits() == 32) { + *(float *)p = (float)v; + } else { + *(float16_t *)p = float16_t(v); + } + } else if (t.is_bfloat()) { + *(bfloat16_t *)p = bfloat16_t(v); + } else if (t.is_int()) { + switch (t.bits()) { + case 8: + *(int8_t *)p = saturate(v); + break; + case 16: + *(int16_t *)p = saturate(v); + break; + case 32: + *(int32_t *)p = saturate(v); + break; + default: + *(int64_t *)p = saturate(v); + break; + } + } else if (t.is_uint() || t.is_bool()) { + switch (t.bits()) { + case 1: + case 8: + *(uint8_t *)p = t.bits() == 1 ? (v != 0) : saturate(v); + break; + case 16: + *(uint16_t *)p = saturate(v); + break; + case 32: + *(uint32_t *)p = saturate(v); + break; + default: + *(uint64_t *)p = saturate(v); + break; + } + } else { + _halide_user_error << "ApproximationTesting: unsupported element type " << t << "\n"; + } +} + +inline double load_value(const void *p, Type t) { + if (t.is_float()) { + if (t.bits() == 64) { + return *(const double *)p; + } else if (t.bits() == 32) { + return *(const float *)p; + } + return (double)*(const float16_t *)p; + } else if (t.is_bfloat()) { + return (double)*(const bfloat16_t *)p; + } else if (t.is_int()) { + switch (t.bits()) { + case 8: + return *(const int8_t *)p; + case 16: + return *(const int16_t *)p; + case 32: + return *(const int32_t *)p; + default: + return (double)*(const int64_t *)p; + } + } else if (t.is_uint() || t.is_bool()) { + switch (t.bits()) { + case 1: + case 8: + return *(const uint8_t *)p; + case 16: + return *(const uint16_t *)p; + case 32: + return *(const uint32_t *)p; + default: + return (double)*(const uint64_t *)p; + } + } + _halide_user_error << "ApproximationTesting: unsupported element type " << t << "\n"; + return 0; +} + +inline uint8_t *element_ptr(const Buffer<> &b, const int *pos) { + const halide_buffer_t *raw = b.raw_buffer(); + int64_t offset = 0; + for (int d = 0; d < raw->dimensions; d++) { + offset += (int64_t)(pos[d] - raw->dim[d].min) * raw->dim[d].stride; + } + return raw->host + offset * raw->type.bytes(); +} + +inline double get(const Buffer<> &b, const int *pos) { + return load_value(element_ptr(b, pos), b.type()); +} + +inline int64_t element_count(const Buffer<> &b) { + int64_t n = 1; + for (int d = 0; d < b.dimensions(); d++) { + n *= b.dim(d).extent(); + } + return n; +} + +// Visit every coordinate of `b`, dimension 0 varying fastest. +template +void for_each_coord(const Buffer<> &b, Fn &&fn) { + if (element_count(b) == 0) { + return; + } + std::vector pos(b.dimensions()); + for (int d = 0; d < b.dimensions(); d++) { + pos[d] = b.dim(d).min(); + } + while (true) { + fn((const int *)pos.data()); + int d = 0; + for (; d < b.dimensions(); d++) { + if (++pos[d] < b.dim(d).min() + b.dim(d).extent()) { + break; + } + pos[d] = b.dim(d).min(); + } + if (d == b.dimensions()) { + return; + } + } +} + +inline std::string coord_string(const std::vector &c) { + std::ostringstream s; + s << "("; + for (size_t i = 0; i < c.size(); i++) { + s << (i ? ", " : "") << c[i]; + } + s << ")"; + return s.str(); +} + +inline std::string number_string(double v, int precision = 9) { + std::ostringstream s; + s.precision(precision); + s << v; + return s.str(); +} + +} // namespace Detail + +// --------------------------------------------------------------------------- +// Distributions +// --------------------------------------------------------------------------- + +/** A description of how to fill a buffer with values. Distributions are + * plain values; nothing is drawn until generate() is called with a seed. + * + * Elements are generated in memory order over the flattened buffer + * (dimension 0 varies fastest), and a *block* is a run of that many + * consecutive elements in this order -- so with dimension 0 the within-block + * index (a (within, block) layout), blocks of size `block` line up with + * dimension 0's extent. + * + * All draws use Rng, and use only integer arithmetic and IEEE +, -, *, so a + * given (distribution, type, extents, seed) yields the same bytes on every + * platform. (normal() is therefore an Irwin-Hall approximation, not + * Box-Muller, which would depend on the platform's libm.) Values are computed + * as doubles and then converted to the element type, saturating for + * integers; 64-bit integer endpoints are therefore only accurate to a + * double's 53 bits. */ +class Distribution { +public: + /** Uniform floating-point values in [lo, hi). */ + static Distribution uniform(double lo, double hi) { + Distribution d(Kind::Uniform); + d.a_ = lo; + d.b_ = hi; + return d; + } + + /** Uniform integers in [lo, hi], both inclusive. */ + static Distribution uniform_int(int64_t lo, int64_t hi) { + Distribution d(Kind::UniformInt); + d.a_ = (double)lo; + d.b_ = (double)hi; + return d; + } + + /** Approximately normal: the sum of twelve uniforms, centered and + * scaled, so its support is mean +- 6 * stddev. */ + static Distribution normal(double mean, double stddev) { + Distribution d(Kind::Normal); + d.a_ = mean; + d.b_ = stddev; + return d; + } + + static Distribution constant(double v) { + Distribution d(Kind::Constant); + d.a_ = v; + return d; + } + + static Distribution zeros() { + return constant(0); + } + + /** Every block of `block` elements is one constant value, drawn once per + * block from `base`. */ + static Distribution blockwise_constant(int block, Distribution base) { + _halide_user_assert(block > 0) << "blockwise_constant: block must be positive\n"; + Distribution d(Kind::Blockwise); + d.block_ = block; + d.parts_ = {std::move(base)}; + return d; + } + + /** Values from `base`, except that one randomly chosen element of every + * block of `block` elements (with a random sign) is `magnitude`. */ + static Distribution outliers(Distribution base, int block, double magnitude) { + _halide_user_assert(block > 0) << "outliers: block must be positive\n"; + Distribution d(Kind::Outliers); + d.block_ = block; + d.a_ = magnitude; + d.parts_ = {std::move(base)}; + return d; + } + + /** Each element is one of the extremes of `type`: for integers its + * minimum and maximum; for floats its lowest and highest finite values + * and +-its smallest normal value. */ + static Distribution extremes(Type type) { + Distribution d(Kind::Extremes); + d.type_ = type; + return d; + } + + /** Each element is one of +-0, +-denormals, +-the smallest normal, +-inf + * or NaN. Floating-point types only. Deliberately never part of any other + * distribution: use it only when the unit under test is meant to cope. */ + static Distribution special_floats() { + return Distribution(Kind::SpecialFloats); + } + + /** Draw each run of `block` consecutive elements entirely from one + * component, chosen uniformly at random (so the default, block = 1, mixes + * per element; a larger block makes whole blocks homogeneous, which is + * what block quantizers are sensitive to). A component generates its + * segment on its own: its own block structure restarts at the segment's + * start. */ + static Distribution mixture(std::vector parts, int block = 1) { + _halide_user_assert(!parts.empty() && block > 0) << "mixture: needs components and a positive block\n"; + Distribution d(Kind::Mixture); + d.block_ = block; + d.parts_ = std::move(parts); + return d; + } + + /** A distribution respecting a port's declared value range (used as the + * default generator for a root's inputs): uniform over the range if the + * port declares one; otherwise over the whole range of an integer type, + * or normal(0, 1) for a floating-point type (whose full range is + * useless for exercising a quantizer). `type` must be the type inputs + * will be generated in. */ + static Distribution from_port(const ApproximationPort &port, Type type) { + if (Detail::is_floating(type)) { + return port.range ? uniform(port.range->lo, port.range->hi) : normal(0, 1); + } + if (port.range) { + return uniform_int((int64_t)std::ceil(port.range->lo), (int64_t)std::floor(port.range->hi)); + } + auto limits = Detail::type_limits(type); + return uniform_int((int64_t)limits.first, (int64_t)limits.second); + } + + /** A short description, e.g. `uniform(-1, 1)`, used as provenance. */ + std::string to_string() const { + using Detail::number_string; + switch (kind_) { + case Kind::Uniform: + return "uniform(" + number_string(a_) + ", " + number_string(b_) + ")"; + case Kind::UniformInt: + return "uniform_int(" + number_string(a_) + ", " + number_string(b_) + ")"; + case Kind::Normal: + return "normal(" + number_string(a_) + ", " + number_string(b_) + ")"; + case Kind::Constant: + return a_ == 0 && !std::signbit(a_) ? "zeros()" : "constant(" + number_string(a_) + ")"; + case Kind::Blockwise: + return "blockwise_constant(" + std::to_string(block_) + ", " + parts_[0].to_string() + ")"; + case Kind::Outliers: + return "outliers(" + parts_[0].to_string() + ", " + std::to_string(block_) + ", " + + number_string(a_) + ")"; + case Kind::Extremes: { + std::ostringstream s; + s << "extremes(" << type_ << ")"; + return s.str(); + } + case Kind::SpecialFloats: + return "special_floats()"; + case Kind::Mixture: { + std::string s = "mixture({"; + for (size_t i = 0; i < parts_.size(); i++) { + s += (i ? ", " : "") + parts_[i].to_string(); + } + return s + "}, " + std::to_string(block_) + ")"; + } + } + return ""; + } + + /** Fill out[0, n) with values drawn for elements of type `type`. */ + void fill(Rng &rng, Type type, double *out, size_t n) const { + switch (kind_) { + case Kind::Constant: + for (size_t i = 0; i < n; i++) { + out[i] = a_; + } + break; + case Kind::Uniform: + for (size_t i = 0; i < n; i++) { + out[i] = a_ + Detail::product(b_ - a_, rng.next_double()); + } + break; + case Kind::UniformInt: + for (size_t i = 0; i < n; i++) { + double v = a_ + std::floor(Detail::product(b_ - a_ + 1, rng.next_double())); + out[i] = v > b_ ? b_ : v; + } + break; + case Kind::Normal: + for (size_t i = 0; i < n; i++) { + double sum = 0; + for (int k = 0; k < 12; k++) { + sum += rng.next_double(); + } + out[i] = a_ + Detail::product(b_, sum - 6.0); + } + break; + case Kind::Blockwise: + for (size_t start = 0; start < n; start += block_) { + double v; + parts_[0].fill(rng, type, &v, 1); + for (size_t i = start; i < std::min(n, start + (size_t)block_); i++) { + out[i] = v; + } + } + break; + case Kind::Outliers: + parts_[0].fill(rng, type, out, n); + for (size_t start = 0; start < n; start += block_) { + size_t len = std::min(n - start, (size_t)block_); + size_t pos = start + rng.next() % len; + out[pos] = (rng.next() & 1) ? -a_ : a_; + } + break; + case Kind::Extremes: { + std::vector values; + if (Detail::is_floating(type_)) { + Detail::FloatInfo info = Detail::float_info(type_); + values = {-info.max, info.max, info.min_normal, -info.min_normal}; + } else { + auto limits = Detail::type_limits(type_); + values = {limits.first, limits.second}; + } + for (size_t i = 0; i < n; i++) { + out[i] = values[rng.next() % values.size()]; + } + break; + } + case Kind::SpecialFloats: { + _halide_user_assert(Detail::is_floating(type)) << "special_floats() requires a floating-point type\n"; + Detail::FloatInfo info = Detail::float_info(type); + const double inf = std::numeric_limits::infinity(); + const std::vector values = { + 0.0, -0.0, info.denorm_min, -info.denorm_min, + info.min_normal - info.denorm_min, -(info.min_normal - info.denorm_min), + info.min_normal, -info.min_normal, inf, -inf, + std::numeric_limits::quiet_NaN()}; + for (size_t i = 0; i < n; i++) { + out[i] = values[rng.next() % values.size()]; + } + break; + } + case Kind::Mixture: + for (size_t start = 0; start < n; start += block_) { + size_t len = std::min(n - start, (size_t)block_); + parts_[rng.next() % parts_.size()].fill(rng, type, out + start, len); + } + break; + } + } + +private: + enum class Kind { Constant, + Uniform, + UniformInt, + Normal, + Blockwise, + Outliers, + Extremes, + SpecialFloats, + Mixture }; + + explicit Distribution(Kind kind) + : kind_(kind), type_(Float(32)) { + } + + Kind kind_; + double a_ = 0, b_ = 0; + int block_ = 1; + Type type_; + std::vector parts_; +}; + +/** Generate a buffer of element type `type` and the given extents (min 0) + * from `dist`, deterministically from `seed`. */ +inline Buffer<> generate(const Distribution &dist, Type type, std::vector extents, uint64_t seed) { + Buffer<> buf(type, extents); + std::vector values((size_t)Detail::element_count(buf)); + Rng rng(seed); + dist.fill(rng, type, values.data(), values.size()); + size_t i = 0; + Detail::for_each_coord(buf, [&](const int *pos) { + Detail::store_value(Detail::element_ptr(buf, pos), type, values[i++]); + }); + return buf; +} + +template +Buffer generate(const Distribution &dist, std::vector extents, uint64_t seed) { + return generate(dist, type_of(), std::move(extents), seed).template as(); +} + +/** How to generate one input of an approximation. */ +struct InputSpec { + Type type; + std::vector extents; + Distribution dist; +}; + +// --------------------------------------------------------------------------- +// Running a round trip +// --------------------------------------------------------------------------- + +namespace Detail { + +struct Capture { + Type type; + int dimensions = 0; + std::map, std::vector> data; +}; + +struct CaptureContext : JITUserContext { + std::map *captures = nullptr; +}; + +inline int32_t capture_trace(JITUserContext *ctx, const halide_trace_event_t *e) { + if (e->event != halide_trace_store) { + return 0; + } + auto &captures = *static_cast(ctx)->captures; + auto it = captures.find(e->func); + if (it == captures.end()) { + return 0; + } + const int lanes = e->lanes, dims = e->dimensions, bytes = e->type.bytes(); + for (int l = 0; l < lanes; l++) { + std::vector coord(dims); + for (int d = 0; d < dims; d++) { + coord[d] = e->coordinates[d * lanes + l]; + } + const uint8_t *src = (const uint8_t *)e->value + (size_t)l * bytes; + it->second.data[coord].assign(src, src + bytes); + } + return 0; +} + +// Assemble captured stores into a buffer covering their bounding box. +inline Buffer<> to_buffer(const Capture &c) { + const int dims = c.dimensions; + std::vector lo(dims, std::numeric_limits::max()), hi(dims, std::numeric_limits::min()); + for (const auto &[coord, bytes] : c.data) { + for (int d = 0; d < dims; d++) { + lo[d] = std::min(lo[d], coord[d]); + hi[d] = std::max(hi[d], coord[d]); + } + } + std::vector extents(dims, 0); + if (!c.data.empty()) { + for (int d = 0; d < dims; d++) { + extents[d] = hi[d] - lo[d] + 1; + } + } + Buffer<> buf(c.type, extents); + if (!c.data.empty()) { + buf.set_min(lo); + } + std::memset(buf.raw_buffer()->host, 0, (size_t)element_count(buf) * c.type.bytes()); + for (const auto &[coord, bytes] : c.data) { + std::memcpy(element_ptr(buf, coord.data()), bytes.data(), bytes.size()); + } + return buf; +} + +inline Func wrap_buffer(const Buffer<> &b, const std::string &name) { + std::vector vars; + std::vector args; + for (int d = 0; d < b.dimensions(); d++) { + vars.emplace_back("tv" + std::to_string(d)); + args.emplace_back(vars.back()); + } + Func f(name); + f(vars) = b(args); + return f; +} + +inline void collect_trace_funcs(const ApproximationTraceNode &node, std::vector &out) { + for (const ApproximationTraceNode &child : node.children) { + collect_trace_funcs(child, out); + } + out.insert(out.end(), node.inputs.begin(), node.inputs.end()); + out.insert(out.end(), node.ports.begin(), node.ports.end()); +} + +inline const ApproximationTraceNode *find_node(const ApproximationTraceNode &node, const Approximation &stage, + int &count) { + const ApproximationTraceNode *found = nullptr; + if (node.stage.same_as(stage)) { + found = &node; + count++; + } + for (const ApproximationTraceNode &child : node.children) { + if (const ApproximationTraceNode *f = find_node(child, stage, count)) { + found = f; + } + } + return found; +} + +// The outcome of running decode(encode(inputs)) with every stage boundary +// captured. +struct RoundTripRun { + EncodeResult enc; + DecodeResult dec; + std::vector> inputs, encoded, decoded; + // Every captured stage-boundary Func's values, by Func name. + std::map> captured; + std::map input_wrappers; + Buffer<> bound; + bool has_bound = false; + + // Could the values of a Func from the encode side be read back? Tuple-valued, + // struct-typed and vector-typed Funcs cannot. + bool can_read(const Func &f) const { + return input_wrappers.count(f.name()) || captured.count(f.name()); + } + + // Why can_read() is false. + static std::string unreadable_reason(const Func &f) { + if (f.outputs() > 1) { + return "Tuple-valued"; + } + return f.types()[0].is_struct() ? "struct-typed" : "not scalar"; + } + + // The realized values of a Func from the encode side. + Buffer<> value_of(const Func &f) const { + auto w = input_wrappers.find(f.name()); + if (w != input_wrappers.end()) { + return inputs[w->second]; + } + auto c = captured.find(f.name()); + _halide_user_assert(c != captured.end()) + << "ApproximationTesting: the values of '" << f.name() << "' were not captured\n"; + return c->second; + } +}; + +inline std::string next_name(const char *base) { + static int counter = 0; + return std::string(base) + "_" + std::to_string(counter++); +} + +// Run a's round trip on `inputs`. Only the encoded Funcs are captured unless +// `capture_all`, which captures every stage's inputs and outputs. +inline RoundTripRun run_round_trip(const Approximation &a, const std::vector> &inputs, + const std::vector &input_names, bool capture_all) { + _halide_user_assert(a.defined() && !inputs.empty()) << "ApproximationTesting: an approximation and inputs are required\n"; + RoundTripRun run; + run.inputs = inputs; + + std::vector in_funcs; + ApproximationPorts ports; + for (size_t i = 0; i < inputs.size(); i++) { + in_funcs.push_back(wrap_buffer(inputs[i], next_name("approx_test_input"))); + run.input_wrappers[in_funcs.back().name()] = (int)i; + if (i < input_names.size()) { + ports.emplace_back(input_names[i]); + } + } + if (ports.size() != inputs.size()) { + ports.clear(); + } + + run.enc = a.encode(in_funcs, ports); + run.dec = a.decode(run.enc.encoded, ports); + _halide_user_assert(run.dec.decoded.size() == inputs.size()) + << "ApproximationTesting: '" << a.label() << "' decoded " << run.dec.decoded.size() + << " Funcs from " << inputs.size() << " inputs\n"; + + // Boundary Funcs to read back. + std::vector wanted = run.enc.encoded; + if (capture_all) { + collect_trace_funcs(run.enc.trace, wanted); + } + std::map captures; + for (const Func &f : wanted) { + if (f.defined() && f.outputs() == 1 && f.types()[0].is_scalar() && !f.types()[0].is_struct() && + !run.input_wrappers.count(f.name()) && + captures.emplace(f.name(), Capture{f.types()[0], f.dimensions(), {}}).second) { + Func(f).compute_root().trace_stores(); + } + } + for (const auto *funcs : {&run.enc.intermediates, &run.dec.intermediates}) { + for (const Func &f : *funcs) { + if (f.has_update_definition()) { + Func(f).compute_root(); + } + } + } + + std::vector outputs = run.dec.decoded; + std::vector> out_buffers; + for (size_t i = 0; i < inputs.size(); i++) { + const Func &d = run.dec.decoded[i]; + _halide_user_assert(d.outputs() == 1 && d.dimensions() == inputs[i].dimensions()) + << "ApproximationTesting: decoded Func '" << d.name() << "' does not match input " << i << "\n"; + std::vector extents, mins; + for (int k = 0; k < inputs[i].dimensions(); k++) { + extents.push_back(inputs[i].dim(k).extent()); + mins.push_back(inputs[i].dim(k).min()); + } + out_buffers.emplace_back(d.types()[0], extents); + out_buffers.back().set_min(mins); + } + Func bound = a.error_bound(in_funcs, run.enc.encoded); + if (bound.defined()) { + _halide_user_assert(bound.dimensions() == inputs[0].dimensions() && bound.outputs() == 1) + << "ApproximationTesting: the declared error bound must be a single-valued Func over the " + << "first input's dimensions\n"; + std::vector vars; + std::vector args; + for (int k = 0; k < bound.dimensions(); k++) { + vars.emplace_back("bv" + std::to_string(k)); + args.emplace_back(vars.back()); + } + Func as_double("approximation_bound_double"); + as_double(vars) = cast(bound(args)); + outputs.push_back(as_double); + std::vector extents, mins; + for (int k = 0; k < inputs[0].dimensions(); k++) { + extents.push_back(inputs[0].dim(k).extent()); + mins.push_back(inputs[0].dim(k).min()); + } + out_buffers.emplace_back(Float(64), extents); + out_buffers.back().set_min(mins); + run.has_bound = true; + } + + CaptureContext ctx; + ctx.captures = &captures; + ctx.handlers.custom_trace = &capture_trace; + Pipeline(outputs).realize(&ctx, Realization(out_buffers)); + + for (const auto &[name, c] : captures) { + run.captured[name] = to_buffer(c); + } + for (size_t i = 0; i < inputs.size(); i++) { + run.decoded.push_back(out_buffers[i]); + } + if (run.has_bound) { + run.bound = out_buffers.back(); + } + // Encoded Funcs that cannot be read back (e.g. struct-typed) are left as undefined Buffers. + for (const Func &f : run.enc.encoded) { + run.encoded.push_back(run.can_read(f) ? run.value_of(f) : Buffer<>()); + } + return run; +} + +// Do two same-shaped values agree exactly? Same-typed buffers are compared +// bitwise; otherwise as doubles. +inline bool same_value(const Buffer<> &a, const Buffer<> &b, const int *pos, bool bitwise) { + if (bitwise && a.type() == b.type()) { + return std::memcmp(element_ptr(a, pos), element_ptr(b, pos), a.type().bytes()) == 0; + } + double x = get(a, pos), y = get(b, pos); + return x == y || (std::isnan(x) && std::isnan(y)); +} + +} // namespace Detail + +// --------------------------------------------------------------------------- +// Round-trip verification +// --------------------------------------------------------------------------- + +/** Error statistics for decode(encode(x)) against x, over all inputs and + * elements. Errors are absolute differences, computed in double; two equal + * values (including two NaNs or two infinities of the same sign) differ by + * zero, and any other difference involving a NaN or infinity is infinite. */ +struct RoundTripReport { + /** The number of values compared. */ + int64_t count = 0; + double max_abs_error = 0; + /** max over values of |y - x| / max(|x|, RoundTripOptions::relative_floor). */ + double max_rel_error = 0; + /** sqrt of the mean squared error. */ + double rmse = 0; + /** The number of values that came back equal (==, or both NaN). */ + int64_t exact_count = 0; + /** Where the largest absolute error is, in which input, and the input + * and decoded values there. */ + std::vector worst_abs_coord; + int worst_input_index = 0; + double worst_input = 0, worst_output = 0; + /** Whether the approximation declares an error bound (Approximation:: + * error_bound()), and how many values of the first input exceed it. */ + bool bound_declared = false; + int64_t bound_violations = 0; + std::vector first_violation_coord; + /** Provenance, for reproducing the run. */ + uint64_t seed = 0; + std::string distribution; +}; + +struct RoundTripOptions { + /** The floor of the denominator of the relative error. If negative, it + * is 1 for integer inputs and the smallest normal float32 otherwise. */ + double relative_floor = -1; +}; + +inline std::ostream &operator<<(std::ostream &s, const RoundTripReport &r) { + using Detail::coord_string; + s << "RoundTripReport: " << r.count << " values, distribution " << r.distribution << ", seed " << r.seed << "\n"; + s << " max abs error " << r.max_abs_error << " at " << coord_string(r.worst_abs_coord) << " of input " + << r.worst_input_index << " (" << r.worst_input << " -> " << r.worst_output << ")\n"; + s << " max rel error " << r.max_rel_error << ", rmse " << r.rmse << ", exact " << r.exact_count << "/" << r.count + << "\n"; + if (r.bound_declared) { + s << " declared bound: " << r.bound_violations << " violation(s)"; + if (r.bound_violations) { + s << ", first at " << coord_string(r.first_violation_coord); + } + s << "\n"; + } else { + s << " declared bound: none\n"; + } + return s; +} + +/** Run `a` on `inputs` (one Buffer per input port, all of which the round + * trip must reproduce) and compare decode(encode(inputs)) with the inputs. + * The declared error bound, if any, is compared against the first input. + * `seed` and `distribution` are recorded in the report as provenance only. */ +inline RoundTripReport verify_round_trip(const Approximation &a, const std::vector> &inputs, + const RoundTripOptions &options = {}, uint64_t seed = 0, + const std::string &distribution = "") { + Detail::RoundTripRun run = Detail::run_round_trip(a, inputs, {}, false); + RoundTripReport report; + report.seed = seed; + report.distribution = distribution; + report.bound_declared = run.has_bound; + double sum_sq = 0; + for (size_t i = 0; i < inputs.size(); i++) { + const Type type = inputs[i].type(); + const double floor = options.relative_floor >= 0 ? options.relative_floor : + Detail::is_integral(type) ? 1.0 : + 1.1754943508222875e-38; + Detail::for_each_coord(inputs[i], [&](const int *pos) { + const double x = Detail::get(inputs[i], pos), y = Detail::get(run.decoded[i], pos); + double err; + if (x == y || (std::isnan(x) && std::isnan(y))) { + err = 0; + report.exact_count++; + } else { + err = std::fabs(y - x); + if (std::isnan(err)) { + err = std::numeric_limits::infinity(); + } + } + if (report.count == 0 || err > report.max_abs_error) { + report.max_abs_error = err; + report.worst_abs_coord.assign(pos, pos + inputs[i].dimensions()); + report.worst_input_index = (int)i; + report.worst_input = x; + report.worst_output = y; + } + report.max_rel_error = std::max(report.max_rel_error, err / std::max(std::fabs(x), floor)); + sum_sq += err * err; + report.count++; + if (i == 0 && run.has_bound && !(err <= Detail::get(run.bound, pos))) { + if (report.bound_violations++ == 0) { + report.first_violation_coord.assign(pos, pos + inputs[i].dimensions()); + } + } + }); + } + report.rmse = report.count ? std::sqrt(sum_sq / (double)report.count) : 0; + return report; +} + +inline RoundTripReport verify_round_trip(const Approximation &a, const Buffer<> &input, + const RoundTripOptions &options = {}) { + return verify_round_trip(a, std::vector>{input}, options); +} + +/** Generate one input of type `type` and the given extents from `dist` and + * `seed`, and verify the round trip on it. */ +inline RoundTripReport verify_round_trip(const Approximation &a, const Distribution &dist, Type type, + std::vector extents, uint64_t seed, + const RoundTripOptions &options = {}) { + Buffer<> input = generate(dist, type, std::move(extents), seed); + return verify_round_trip(a, std::vector>{input}, options, seed, dist.to_string()); +} + +// --------------------------------------------------------------------------- +// Properties +// --------------------------------------------------------------------------- + +/** The outcome of one property check on one trial. */ +struct PropertyOutcome { + bool passed = true; + std::string message; + /** The (first, in memory order) failing coordinate, if the failure has + * one. */ + std::vector coord; +}; + +/** What a property is checked against: the approximation under test (the + * root, or a stage for a stage-targeted property), the inputs it received, + * and its realized round trip. */ +struct PropertyContext { + Approximation approximation; + std::vector> inputs; + std::vector input_names; + Detail::RoundTripRun run; + + /** The Buffers the encode side produced (one per encoded Func). An + * encoded Func that is Tuple-valued or struct-typed cannot be read back, + * and its Buffer is undefined; see encoded_readable(). */ + const std::vector> &encoded() const { + return run.encoded; + } + /** Can encoded()[p] be used? */ + bool encoded_readable(size_t p) const { + return p < run.encoded.size() && run.encoded[p].defined(); + } + /** A failure describing encoded Func `p` as unreadable, for properties + * that need its values. */ + PropertyOutcome encoded_unreadable(size_t p) const; + /** The decoded Buffers (one per input). */ + const std::vector> &decoded() const { + return run.decoded; + } + /** The ports of encoded(), including declared value ranges. */ + const ApproximationPorts &encoded_ports() const { + return run.enc.encoded_ports; + } + + /** The round trip run again on decoded(), for properties about + * re-encoding. Computed on first use. */ + const Detail::RoundTripRun &requantized() const { + if (!requantized_) { + requantized_ = Detail::run_round_trip(approximation, run.decoded, input_names, false); + } + return *requantized_; + } + +private: + mutable std::optional requantized_; +}; + +/** A named check of an approximation on realized values. Build one with the + * factory functions below (or from a function of a PropertyContext), and + * retarget it at an inner stage with at(). */ +class Property { +public: + using Check = std::function; + + Property(std::string name, Check check) + : name_(std::move(name)), check_(std::move(check)) { + } + + const std::string &name() const { + return name_; + } + const Check &check() const { + return check_; + } + + /** The stage a stage-targeted property is about, or an undefined handle. */ + const Approximation &stage() const { + return stage_; + } + + /** A copy of this property that check_property() evaluates for `stage` + * alone, on the values that actually arrive at its encode inputs when + * the whole approximation runs on the generated inputs. `stage` must be + * invoked exactly once by the approximation under test (hold on to the + * handle you composed it with). The arriving values are also checked + * against the stage's declared input ranges; see PropertyOptions. */ + Property at(const Approximation &stage) const { + Property p = *this; + p.stage_ = stage; + return p; + } + +private: + std::string name_; + Check check_; + Approximation stage_; +}; + +namespace Detail { + +inline PropertyOutcome fail(const std::string &message, const std::vector &coord = {}) { + PropertyOutcome o; + o.passed = false; + o.message = message; + o.coord = coord; + return o; +} + +} // namespace Detail + +inline PropertyOutcome PropertyContext::encoded_unreadable(size_t p) const { + std::string name = p < run.enc.encoded_ports.size() ? run.enc.encoded_ports[p].name : std::to_string(p); + return Detail::fail("the encoded port '" + name + "' is " + Detail::RoundTripRun::unreadable_reason(run.enc.encoded[p]) + + " and cannot be read back, so this property does not apply to it"); +} + +namespace Detail { + +inline std::string describe_pair(double x, double y) { + return number_string(x, 9) + " -> " + number_string(y, 9); +} + +// Check `test(x, y, in_buffer, pos)` over every input/decoded pair; report the first failure. +template +PropertyOutcome check_pairs(const PropertyContext &c, const char *what, Test &&test) { + for (size_t i = 0; i < c.inputs.size(); i++) { + PropertyOutcome result; + for_each_coord(c.inputs[i], [&](const int *pos) { + if (!result.passed) { + return; + } + std::string problem = test(get(c.inputs[i], pos), get(c.decoded()[i], pos), c.inputs[i], c.decoded()[i], pos); + if (!problem.empty()) { + result = fail("input " + std::to_string(i) + " at " + + coord_string(std::vector(pos, pos + c.inputs[i].dimensions())) + ": " + + what + " (" + problem + ")", + std::vector(pos, pos + c.inputs[i].dimensions())); + } + }); + if (!result.passed) { + return result; + } + } + return {}; +} + +} // namespace Detail + +/** decode(encode(x)) is bit-for-bit x. (-0.0 versus 0.0 counts as a + * difference; NaNs with equal bits do not.) Losslessness typically holds + * only for inputs within the declared input ranges. */ +inline Property lossless() { + return Property("lossless", [](const PropertyContext &c) { + PropertyOutcome o; + for (size_t i = 0; i < c.inputs.size() && o.passed; i++) { + Detail::for_each_coord(c.inputs[i], [&](const int *pos) { + if (o.passed && !Detail::same_value(c.inputs[i], c.decoded()[i], pos, true)) { + std::vector coord(pos, pos + c.inputs[i].dimensions()); + o = Detail::fail("input " + std::to_string(i) + " at " + Detail::coord_string(coord) + + ": round trip changed the value (" + + Detail::describe_pair(Detail::get(c.inputs[i], pos), Detail::get(c.decoded()[i], pos)) + ")", + coord); + } + }); + } + return o; + }); +} + +/** |decode(encode(x)) - x| <= abs everywhere. Values that come back equal + * always pass; a NaN or infinity that does not come back equal always fails. */ +inline Property bounded_error(double abs) { + return Property("bounded_error(" + Detail::number_string(abs) + ")", [abs](const PropertyContext &c) { + return Detail::check_pairs(c, "error exceeds the bound", [&](double x, double y, const Buffer<> &, const Buffer<> &, const int *) { + if (x == y || (std::isnan(x) && std::isnan(y))) { + return std::string(); + } + double err = std::fabs(y - x); + return err <= abs ? std::string() : Detail::describe_pair(x, y) + ", error " + Detail::number_string(err); + }); + }); +} + +/** |decode(encode(x)) - x| <= the unit's declared error_bound(), per element + * of the first input. Fails if the approximation declares no bound. */ +inline Property within_declared_bound() { + return Property("within_declared_bound", [](const PropertyContext &c) { + if (!c.run.has_bound) { + return Detail::fail("'" + c.approximation.label() + "' declares no error bound"); + } + PropertyOutcome o; + Detail::for_each_coord(c.inputs[0], [&](const int *pos) { + if (!o.passed) { + return; + } + double x = Detail::get(c.inputs[0], pos), y = Detail::get(c.decoded()[0], pos); + double err = (x == y || (std::isnan(x) && std::isnan(y))) ? 0 : std::fabs(y - x); + double bound = Detail::get(c.run.bound, pos); + if (!(err <= bound)) { + std::vector coord(pos, pos + c.inputs[0].dimensions()); + o = Detail::fail("at " + Detail::coord_string(coord) + ": error " + Detail::number_string(err) + + " exceeds the declared bound " + Detail::number_string(bound) + " (" + + Detail::describe_pair(x, y) + ")", + coord); + } + }); + return o; + }); +} + +/** Re-encoding a decoded value reproduces the encoding: encode(decode(e)) + * == e, for e = encode(x), elementwise on every encoded Func. Integer + * encodings must match exactly. Floating-point encodings (e.g. a + * quantizer's scale, recomputed as `qmax * scale / qmax`) may differ in the + * last bits, so they match if within `float_rel_tol` relative to the larger + * magnitude. */ +inline Property idempotent_requantize(double float_rel_tol = 1e-6) { + return Property("idempotent_requantize", [float_rel_tol](const PropertyContext &c) { + for (size_t p = 0; p < c.encoded().size(); p++) { + if (!c.encoded_readable(p)) { + return c.encoded_unreadable(p); + } + } + const Detail::RoundTripRun &again = c.requantized(); + if (again.encoded.size() != c.encoded().size()) { + return Detail::fail("re-encoding produced " + std::to_string(again.encoded.size()) + " Funcs instead of " + + std::to_string(c.encoded().size())); + } + for (size_t p = 0; p < c.encoded().size(); p++) { + const Buffer<> &e1 = c.encoded()[p], &e2 = again.encoded[p]; + if (!e2.defined() || e1.dimensions() != e2.dimensions() || e1.type() != e2.type()) { + return Detail::fail("encoded Func " + std::to_string(p) + " changed shape or type on re-encoding"); + } + PropertyOutcome o; + Detail::for_each_coord(e1, [&](const int *pos) { + if (!o.passed) { + return; + } + bool ok; + double x = Detail::get(e1, pos), y = Detail::get(e2, pos); + if (Detail::is_floating(e1.type())) { + ok = x == y || (std::isnan(x) && std::isnan(y)) || + std::fabs(x - y) <= float_rel_tol * std::max(std::fabs(x), std::fabs(y)); + } else { + ok = x == y; + } + if (!ok) { + std::vector coord(pos, pos + e1.dimensions()); + o = Detail::fail("encoded Func " + std::to_string(p) + " at " + Detail::coord_string(coord) + + " changed on re-encoding (" + Detail::describe_pair(x, y) + ")", + coord); + } + }); + if (!o.passed) { + return o; + } + } + return PropertyOutcome(); + }); +} + +/** Inputs that are zero decode to zero (== 0; -0.0 counts as zero). */ +inline Property zero_preserving() { + return Property("zero_preserving", [](const PropertyContext &c) { + return Detail::check_pairs(c, "zero did not round-trip to zero", [](double x, double y, const Buffer<> &, const Buffer<> &, const int *) { + return x == 0 && y != 0 ? Detail::describe_pair(x, y) : std::string(); + }); + }); +} + +/** The round trip never flips a sign: a positive input decodes to a value >= 0 + * and a negative one to a value <= 0 (so flushing to zero is allowed). NaNs + * are ignored. */ +inline Property sign_preserving() { + return Property("sign_preserving", [](const PropertyContext &c) { + return Detail::check_pairs(c, "sign flipped", [](double x, double y, const Buffer<> &, const Buffer<> &, const int *) { + return (x > 0 && y < 0) || (x < 0 && y > 0) ? Detail::describe_pair(x, y) : std::string(); + }); + }); +} + +/** Every value of every encoded Func lies in that port's declared output + * range (a guarantee, see ApproximationPort::range). Ports with no declared + * range are not checked. */ +inline Property outputs_within_declared_ranges() { + return Property("outputs_within_declared_ranges", [](const PropertyContext &c) { + for (size_t p = 0; p < c.encoded().size() && p < c.encoded_ports().size(); p++) { + const ApproximationPort &port = c.encoded_ports()[p]; + if (!port.range) { + continue; + } + if (!c.encoded_readable(p)) { + return c.encoded_unreadable(p); + } + PropertyOutcome o; + Detail::for_each_coord(c.encoded()[p], [&](const int *pos) { + double v = Detail::get(c.encoded()[p], pos); + if (o.passed && !port.range->contains(v)) { + std::vector coord(pos, pos + c.encoded()[p].dimensions()); + o = Detail::fail("output '" + port.name + "' at " + Detail::coord_string(coord) + " is " + + Detail::number_string(v) + ", outside the declared range [" + + Detail::number_string(port.range->lo) + ", " + + Detail::number_string(port.range->hi) + "]", + coord); + } + }); + if (!o.passed) { + return o; + } + } + return PropertyOutcome(); + }); +} + +// --------------------------------------------------------------------------- +// check_property +// --------------------------------------------------------------------------- + +/** The outcome of check_property(). On failure, `failing_seed` is the seed of + * the failing trial: `check_property(..., trials = 1, seed = failing_seed)` + * reproduces it exactly. */ +struct PropertyResult { + bool passed = true; + std::string property; + /** The label of the stage checked (the root's label if not stage-targeted). */ + std::string stage_label; + /** The number of trials run, including the failing one. */ + int trials_run = 0; + uint64_t failing_seed = 0; + /** The failing coordinate in the values checked (empty if the failure + * has none). */ + std::vector failing_coord; + std::string message; + /** True if the trial failed because its inputs lay outside the declared + * input ranges of the stage checked, rather than because the property + * itself failed. */ + bool precondition_violated = false; +}; + +inline std::ostream &operator<<(std::ostream &s, const PropertyResult &r) { + s << "property " << r.property << " on '" << r.stage_label << "': "; + if (r.passed) { + return s << "passed (" << r.trials_run << " trial" << (r.trials_run == 1 ? "" : "s") << ")\n"; + } + s << "FAILED in trial " << r.trials_run << " (seed " << r.failing_seed << ")"; + if (!r.failing_coord.empty()) { + s << " at " << Detail::coord_string(r.failing_coord); + } + return s << ": " << r.message << "\n"; +} + +struct PropertyOptions { + /** Preconditions: the values reaching the stage checked (the generated + * inputs, or for `.at(stage)` the arriving values) are compared with that + * stage's declared input ranges. If true (the default), a trial with a + * value outside them is *not run* and fails with + * PropertyResult::precondition_violated set: the property is only claimed + * where its preconditions hold, so a generator (or upstream stage) that + * breaks them is reported as such. If false, the property runs anyway, + * which is how to demonstrate that it fails outside its preconditions. */ + bool check_preconditions = true; +}; + +using InputGenerator = std::function>(uint64_t seed)>; + +namespace Detail { + +// Compare values with declared preconditions; the empty string if they hold. +inline std::string precondition_problem(const ApproximationPorts &required, const std::vector> &values, + std::vector &coord) { + for (size_t i = 0; i < required.size() && i < values.size(); i++) { + if (!required[i].range) { + continue; + } + const ApproximationRange &range = *required[i].range; + std::string problem; + for_each_coord(values[i], [&](const int *pos) { + double v = get(values[i], pos); + if (problem.empty() && !range.contains(v)) { + coord.assign(pos, pos + values[i].dimensions()); + problem = "precondition not met: input '" + required[i].name + "' is " + number_string(v) + + " at " + coord_string(coord) + ", outside the declared range [" + + number_string(range.lo) + ", " + number_string(range.hi) + "]"; + } + }); + if (!problem.empty()) { + return problem; + } + } + return ""; +} + +} // namespace Detail + +/** Check `prop` on `trials` trials, each on freshly generated inputs (from + * `gen`, given the trial's seed, which is derive_seed(seed, trial)), and stop + * at the first failure. + * + * For a plain property, the property is about `a` on the generated inputs. + * For prop.at(stage), `a` is run on the generated inputs, the values that + * arrive at `stage`'s encode inputs are read back, and the property is + * evaluated for `stage` alone (encode, then decode) on those values. Either + * way the values that go into the checked unit are first compared with its + * declared input ranges (see PropertyOptions). */ +inline PropertyResult check_property(const Approximation &a, const Property &prop, const InputGenerator &gen, + int trials, uint64_t seed, const PropertyOptions &options = {}) { + PropertyResult result; + result.property = prop.name(); + const Approximation &target = prop.stage().defined() ? prop.stage() : a; + result.stage_label = target.label(); + + for (int trial = 0; trial < trials; trial++) { + const uint64_t trial_seed = Halide::ApproximationTesting::derive_seed(seed, trial); + result.trials_run = trial + 1; + std::vector> inputs = gen(trial_seed); + + PropertyContext ctx; + std::vector> stage_inputs; + std::vector stage_names; + if (prop.stage().defined()) { + Detail::RoundTripRun root = Detail::run_round_trip(a, inputs, {}, true); + int count = 0; + const ApproximationTraceNode *node = Detail::find_node(root.enc.trace, prop.stage(), count); + _halide_user_assert(count == 1) << "check_property: the stage '" << prop.stage().label() << "' was invoked " + << count << " times by the encode of '" << a.label() << "' (expected once)\n"; + for (const Func &f : node->inputs) { + if (!root.can_read(f)) { + result.passed = false; + result.failing_seed = trial_seed; + result.message = "the input '" + f.name() + "' of stage '" + prop.stage().label() + "' is " + + Detail::RoundTripRun::unreadable_reason(f) + + " and cannot be read back, so this property cannot be checked at that stage"; + return result; + } + stage_inputs.push_back(root.value_of(f)); + } + stage_names = node->input_names; + } else { + stage_inputs = inputs; + } + + if (options.check_preconditions) { + ApproximationPorts context; + for (const std::string &name : stage_names) { + context.emplace_back(name); + } + ApproximationSignature sig = target.signature(context); + std::vector coord; + std::string problem = sig.known ? Detail::precondition_problem(sig.inputs, stage_inputs, coord) : ""; + if (!problem.empty()) { + result.passed = false; + result.failing_seed = trial_seed; + result.failing_coord = coord; + result.message = problem; + result.precondition_violated = true; + return result; + } + } + + ctx.approximation = target; + ctx.inputs = stage_inputs; + ctx.input_names = stage_names; + ctx.run = Detail::run_round_trip(target, stage_inputs, stage_names, false); + PropertyOutcome outcome = prop.check()(ctx); + if (!outcome.passed) { + result.passed = false; + result.failing_seed = trial_seed; + result.failing_coord = outcome.coord; + result.message = outcome.message; + return result; + } + } + return result; +} + +/** As above, generating one buffer per InputSpec (input k is seeded with + * derive_seed(trial_seed, k)). */ +inline PropertyResult check_property(const Approximation &a, const Property &prop, + const std::vector &specs, int trials = 8, uint64_t seed = 0, + const PropertyOptions &options = {}) { + InputGenerator gen = [specs](uint64_t s) { + std::vector> inputs; + for (size_t k = 0; k < specs.size(); k++) { + inputs.push_back(generate(specs[k].dist, specs[k].type, specs[k].extents, + Halide::ApproximationTesting::derive_seed(s, k))); + } + return inputs; + }; + return check_property(a, prop, gen, trials, seed, options); +} + +/** As above, for a single input of type `type`, drawn from `dist`. */ +inline PropertyResult check_property(const Approximation &a, const Property &prop, const Distribution &dist, + Type type, std::vector extents, int trials = 8, uint64_t seed = 0, + const PropertyOptions &options = {}) { + return check_property(a, prop, std::vector{{type, std::move(extents), dist}}, trials, seed, options); +} + +/** The default generator: one input per declared root input port, all with + * the given extents, each drawn from Distribution::from_port() -- uniform + * within the port's declared range, so the property is exercised exactly + * where its preconditions hold. An input port without a declared type is + * generated as float32. The approximation must have a known signature. */ +inline PropertyResult check_property(const Approximation &a, const Property &prop, std::vector extents, + int trials = 8, uint64_t seed = 0, const PropertyOptions &options = {}) { + ApproximationSignature sig = a.signature(); + _halide_user_assert(sig.known && !sig.inputs.empty()) + << "check_property: '" << a.label() << "' has no declared input ports to generate from; pass InputSpecs\n"; + std::vector specs; + for (const ApproximationPort &port : sig.inputs) { + Type type = port.type.value_or(Float(32)); + specs.push_back({type, extents, Distribution::from_port(port, type)}); + } + return check_property(a, prop, specs, trials, seed, options); +} + +} // namespace ApproximationTesting +} // namespace Halide + +#endif diff --git a/tutorial/CMakeLists.txt b/tutorial/CMakeLists.txt index 94c3d9e22be7..aa367621eef2 100644 --- a/tutorial/CMakeLists.txt +++ b/tutorial/CMakeLists.txt @@ -341,7 +341,8 @@ if (TARGET Halide::Mullapudi2016) ) endif () -# Lessons 22-24 +# Lessons 22-25 add_tutorial(lesson_22_jit_performance.cpp) add_tutorial(lesson_23_serialization.cpp WITH_IMAGE_IO) add_tutorial(lesson_24_async.cpp GROUPS multithreaded) +add_tutorial(lesson_25_approximations.cpp) diff --git a/tutorial/lesson_25_approximations.cpp b/tutorial/lesson_25_approximations.cpp new file mode 100644 index 000000000000..75c06be29c99 --- /dev/null +++ b/tutorial/lesson_25_approximations.cpp @@ -0,0 +1,477 @@ +// Halide tutorial lesson 25: Approximations (GGML's Q4_0 weight format) + +// This lesson shows how to describe a lossy, quantized data format with +// Halide's Approximation system, using GGML's Q4_0 block format as the +// running example. By the end we will have: +// - built Q4_0 out of reusable components, +// - spliced it into an existing pipeline (a dot product) with +// Func::approximate_by, +// - split the pipeline into an offline quantizer and an online consumer +// with Pipeline::sever, +// - checked that the bytes we produce are bit-for-bit what GGML produces, +// - and tested the properties the scheme claims. + +// On linux, you can compile and run it like so: +// g++ lesson_25*.cpp -g -I -I -L -lHalide -lpthread -ldl -o lesson_25 -std=c++17 +// LD_LIBRARY_PATH= ./lesson_25 + +// On macOS: +// g++ lesson_25*.cpp -g -I -I -L -lHalide -o lesson_25 -std=c++17 +// DYLD_LIBRARY_PATH= ./lesson_25 + +// Halide.h contains the Approximation machinery, but the property-testing +// helpers live in a separate header-only tool, like halide_image_io.h. That +// header is used with a plain #include of Halide.h. +#include "Halide.h" +#include "halide_approximation_testing.h" + +#include +#include +#include +#include +#include +#include +#include + +using namespace Halide; +using namespace Halide::ApproximationTesting; + +// ---------------------------------------------------------------------------- +// Part 0: The plain C++ reference. +// +// This is a transcription of quantize_row_q4_0_ref and dequantize_row_q4_0 +// from GGML's ggml-quants.c. We will use it at the end to check the +// Approximation bit-for-bit. +// ---------------------------------------------------------------------------- + +constexpr int QK4_0 = 32; + +// assert() disappears in release builds, so use our own. +void require(bool ok, const char *what) { + if (!ok) { + printf("Check failed: %s\n", what); + exit(1); + } +} + +struct block_q4_0 { + uint16_t d; // the fp16 scale, as raw bits + uint8_t qs[QK4_0 / 2]; +}; +static_assert(sizeof(block_q4_0) == 18, "block_q4_0 must be 18 bytes"); + +void reference_quantize(const float *x, block_q4_0 *y, int nblocks) { + for (int i = 0; i < nblocks; i++) { + float amax = 0.0f; + float max = 0.0f; + for (int j = 0; j < QK4_0; j++) { + float v = x[i * QK4_0 + j]; + if (amax < std::fabs(v)) { + amax = std::fabs(v); + max = v; + } + } + const float d = max / -8; + const float id = d != 0.0f ? 1.0f / d : 0.0f; + y[i].d = float16_t(d).to_bits(); + for (int j = 0; j < QK4_0 / 2; j++) { + // Separate statements, so the compiler can't fuse them into an fma. + const float x0 = x[i * QK4_0 + j] * id; + const float x1 = x[i * QK4_0 + QK4_0 / 2 + j] * id; + const uint8_t xi0 = std::min(15, (int8_t)(x0 + 8.5f)); + const uint8_t xi1 = std::min(15, (int8_t)(x1 + 8.5f)); + y[i].qs[j] = xi0 | (xi1 << 4); + } + } +} + +void reference_dequantize(const block_q4_0 *x, float *y, int nblocks) { + for (int i = 0; i < nblocks; i++) { + const float d = (float)float16_t::make_from_bits(x[i].d); + for (int j = 0; j < QK4_0 / 2; j++) { + y[i * QK4_0 + j] = ((x[i].qs[j] & 0x0F) - 8) * d; + y[i * QK4_0 + j + QK4_0 / 2] = ((x[i].qs[j] >> 4) - 8) * d; + } + } +} + +// Deterministic test data. Most blocks are pseudo-random, but the first few +// are tricky on purpose. +std::vector make_weights(int nblocks) { + std::vector w(static_cast(nblocks) * QK4_0); + uint32_t state = 12345; + for (float &v : w) { + state = state * 1664525u + 1013904223u; + v = ((int32_t)(state >> 8) / (float)(1 << 23) - 1.0f) * 3.0f; + } + for (int j = 0; j < QK4_0; j++) { + w[0 * QK4_0 + j] = 0.0f; // all zeros: d == 0 + w[1 * QK4_0 + j] = (j % 2 != 0 ? -1 : 1) * 0.25f * (j % 8); // ties: +/-1.75 repeat + w[2 * QK4_0 + j] = -0.5f - 0.1f * j; // the extreme is negative + w[3 * QK4_0 + j] = 100.0f; // constant block + w[4 * QK4_0 + j] = 1e-20f * (j - 7); // tiny scale + } + return w; +} + +// ---------------------------------------------------------------------------- +// Part 1: What is an Approximation? +// +// A quantized weight format is an *approximate identity* factored in two: +// decode(encode(x)) ~= x +// `encode` is run once, offline, to compress the weights; `decode` runs every +// time the weights are used. An Approximation bundles the pair so that a +// scheme can't get out of sync with itself. +// +// The smallest possible Approximation is a Pointwise unit: an elementwise +// Expr -> Expr function for each direction. +// ---------------------------------------------------------------------------- + +void part1_a_tiny_approximation() { + // Drop the low bit of an integer: lossy, with error at most 1. + Approximation drop_lsb = + Pointwise{"drop_lsb", + [](const Expr &x) { return x >> 1; }, + [](const Expr &x) { return x << 1; }} + .with_error_bound([](const Expr &) { return Expr(1); }); + + Var x("x"); + Func f("f"), consumer("consumer"); + f(x) = x; + consumer(x) = f(x) + 100; + + // approximate_by rewrites `consumer` so that wherever it used f(x), it now + // uses decode(encode(f))(x). Note that it edits the algorithm: unlike a + // schedule directive, this changes the values that the pipeline computes. + f.approximate_by(drop_lsb, {consumer}); + + Buffer out = consumer.realize({8}); + for (int i = 0; i < 8; i++) { + require(out(i) == 100 + (i & ~1), "drop_lsb round trip"); + } +} + +// ---------------------------------------------------------------------------- +// Part 2: Q4_0 out of reusable components. +// +// Q4_0 stores each block of 32 floats as an fp16 scale `d` and sixteen bytes +// of 4-bit codes. We can read off the encode direction as a pipeline of small +// steps, each of which is an Approximation of its own. Compose takes them +// in the order encode runs them, and decode runs them backwards: +// +// BlockReshape{32} flat floats -> (within, block) +// Q4_0Quantizer (within, block) -> int8 codes [-8, 7], fp32 scale +// offset codes [-8, 7] -> nibbles [0, 15] +// PlanarFieldPack{4, 16} nibbles -> 16 bytes, with element j in the low +// half of byte j and element j+16 in the high half +// fp16 fp32 scale -> fp16 scale +// StructLayout {qs, d} -> one 18-byte record +// +// Everything but the quantizer is generic. The quantizer is where Q4_0's +// arithmetic lives, and it is ordinary client code: a struct with an encode +// and a decode method (and, optionally, declarations about itself), which +// converts implicitly to an Approximation. +// ---------------------------------------------------------------------------- + +struct Q4_0Quantizer { + // Encode makes both the codes and the scale, so it takes and returns + // vectors of Funcs. + static std::vector encode(const std::vector &in) { + const Func &blocks = in[0]; + Var j("j"), b("b"); + + // The signed element of largest magnitude in each block. The first + // one wins ties, so the comparison is strict. + RDom r(0, QK4_0); + Func extreme("extreme"); + extreme(b) = Tuple(0.0f, 0.0f); + Expr v = blocks(r, b); + Expr bigger = abs(v) > extreme(b)[0]; + extreme(b) = Tuple(select(bigger, abs(v), extreme(b)[0]), + select(bigger, v, extreme(b)[1])); + + Func scale("scale"), codes("codes"); + scale(b) = extreme(b)[1] / -8.0f; + Expr inv_scale = select(scale(b) != 0.0f, 1.0f / scale(b), 0.0f); + // One add of 8.5f, as in GGML: (int8_t)(x * id + 8.5f). + Expr biased = cast(blocks(j, b) * inv_scale + 8.5f); + codes(j, b) = min(biased, cast(15)) - 8; + return {codes, scale}; + } + + static std::vector decode(const std::vector &encoded) { + const Func &codes = encoded[0], &scale = encoded[1]; + Var j("j"), b("b"); + Func values("values"); + values(j, b) = codes(j, b) * scale(b); + return {values}; + } + + // Optional: the types and dimensions of the ports, and a guarantee about + // the codes. Later stages' preconditions are checked against it. + static ApproximationSignature signature() { + return {{{"blocks", Float(32), 2}}, + {{"codes", Int(8), 2, ApproximationRange(-8, 7)}, + {"scale", Float(32), 1}}}; + } + + // Optional: |decode(encode(x)) - x| is at most one step (the top code is + // clamped: a value that scales to +8 becomes 7), plus slack for rounding. + static Func error_bound(const std::vector & /*inputs*/, const std::vector &encoded) { + Var j("j"), b("b"); + Func bound("bound"); + bound(j, b) = abs(cast(encoded[1](b))) * Expr(1.0001); + return bound; + } +}; + +// Pointwise units are elementwise conversions, written inline. This one +// shifts the signed codes into the unsigned nibbles that PlanarFieldPack packs. +// The declarations are optional: the types are checked, the input range is a +// precondition for being lossless, and the output range is a guarantee. +Pointwise make_offset() { + return Pointwise{"offset", + [](const Expr &x) { return cast(x + 8); }, + [](const Expr &x) { return cast(cast(x) - 8); }} + .with_types(Int(8), UInt(8)) + .with_ranges(ApproximationRange(-8, 7), ApproximationRange(0, 15)) + .with_lossless(); +} + +// Likewise, a cast: the fp32 scale is stored as fp16. +Pointwise make_fp16() { + return Pointwise{"fp16", + [](const Expr &x) { return cast(x); }, + [](const Expr &x) { return cast(x); }} + .with_types(Float(32), Float(16)); +} + +// A scheme is found again later (to schedule it, or to read back one of its +// intermediate values) through the handles of the stages it was built from, so +// we keep the interesting ones around next to the finished scheme. +struct Q4_0 { + Approximation quantize = Q4_0Quantizer{}; + Approximation offset = Approximation(make_offset(), "offset"); + Approximation pack = PlanarFieldPack{4, QK4_0 / 2}; + Approximation fp16 = Approximation(make_fp16(), "fp16"); + + // block_q4_0, as a Halide struct type. Its size is 18 bytes, and it has + // the same layout as the C++ struct above. + static Type block_type() { + return Type::Struct({{"d", Float(16)}, {"qs", UInt(8), 16}}); + } + + // Stages are listed in encode order; decode runs them backwards. Parallel + // routes each named port to its own child, and passes the others through. + // A port keeps its name in both directions, so "codes" and "scale" are the + // right handles on the decode side too. + Approximation scheme = Compose{ + BlockReshape{QK4_0}, + quantize, + Parallel{{"codes", Compose{offset, pack}}, + {"scale", fp16}}, + StructLayout{block_type(), {"qs", "d"}}}; +}; + +// ---------------------------------------------------------------------------- +// Part 3: Using it. +// +// Suppose we already have a pipeline: the dot product of a weight vector with +// an activation vector. It was written against float weights, and we'd like +// to store them as Q4_0. +// ---------------------------------------------------------------------------- + +int main() { + part1_a_tiny_approximation(); + + const int nblocks = 256; + const int N = nblocks * QK4_0; + + // The pipeline as originally written: + ImageParam weights_in(Float(32), 1, "weights_in"), acts_in(Float(32), 1, "acts_in"); + Var k("k"); + Func weights("weights"); + weights(k) = weights_in(k); + RDom r(0, N, "r"); + Func dot("dot"); + dot() = 0.0f; + dot() += weights(r) * acts_in(r); + + // Now approximate the weights, and splice the round trip into `dot`. The + // ApproximationResult describes everything that was created. + Q4_0 q; + Approximation q4_0 = q.scheme; + + // Printing an Approximation shows its structure, with the type and range + // of every port, without running anything: + // + // Compose (values x1) -> (record: struct{d: float16, qs: uint8[16]} x1) + // BlockReshape (values x1) -> (blocks x2) + // Q4_0Quantizer (blocks: float32 x2) -> (codes: int8 x2 in [-8, 7], scale: float32 x1) + // Parallel (codes: int8 x2 in [-8, 7], scale: float32 x1) -> (bytes: uint8 x2, scale: float16 x1) + // Compose (codes: int8 x2 in [-8, 7]) -> (bytes: uint8 x2) + // offset (codes: int8 x2 in [-8, 7]) -> (codes: uint8 x2 in [0, 15]) + // PlanarFieldPack (codes x2 in [0, 15]) -> (bytes: uint8 x2) + // fp16 (scale: float32 x1) -> (scale: float16 x1) + // StructLayout (bytes: uint8 x2, scale: float16 x1) -> (record: struct{d: float16, qs: uint8[16]} x1) + std::cout << q4_0 << "\n"; + ApproximationResult approx = weights.approximate_by(q4_0, {dot}); + + // Printing `approx` (try it!) shows the tree of stages that ran, with the + // Funcs each one produced and their helper Funcs, e.g.: + // + // encode: + // Q4_0Quantizer -> codes=codes, scale=scale + // intermediates: extreme + // ... + + // Anything with an update definition (like the per-block search for the + // largest element, `extreme`) or that is the boundary between two stages is an + // ordinary Func. Here, we compute those at root. The rest are pure, and + // get inlined into their consumers. + for (Func f : approx.intermediates) { + if (f.has_update_definition() || approx.is_stage_port(f)) { + f.compute_root(); + } + } + + // Find one particular Func by the handle of the stage that made it. The + // per-block scale that Q4_0 divides by is an fp32 value, before it is + // rounded to fp16 for storage. + Func fp32_scale = approx.encoded_by(q.quantize, "scale"); + + // ------------------------------------------------------------------------ + // Part 4: Offline and online. + // + // Quantizing is done once, when a model is converted. The dot product is + // done every time it's run. sever severs the pipeline at the + // encoded weights: `split.offline` computes them, and the pipeline we + // gave it now reads them from an ImageParam instead. + // ------------------------------------------------------------------------ + SeverResult split = Pipeline(dot).sever(approx.encoded); + + std::vector weights_data = make_weights(nblocks); + std::vector acts_data(N); + for (int i = 0; i < N; i++) { + acts_data[i] = std::sin(0.01f * i); + } + weights_in.set(Buffer(weights_data.data(), N)); + acts_in.set(Buffer(acts_data.data(), N)); + + // Offline: quantize. The result is one struct per block. + Buffer<> encoded(Q4_0::block_type(), nblocks); + split.offline.realize(encoded); + + // Online: from here on, nothing looks at the fp32 weights. + split.online_inputs[0].set(encoded); + Buffer dot_result = dot.realize(); + Buffer dequantized = approx.replacement.realize({N}); + + // ------------------------------------------------------------------------ + // Part 5: Is it really Q4_0? + // + // Compare against the C++ reference, byte for byte. + // ------------------------------------------------------------------------ + std::vector ref_blocks(nblocks); + reference_quantize(weights_data.data(), ref_blocks.data(), nblocks); + std::vector ref_dequantized(N); + reference_dequantize(ref_blocks.data(), ref_dequantized.data(), nblocks); + + require(encoded.size_in_bytes() == nblocks * sizeof(block_q4_0), "encoded size"); + if (memcmp(encoded.data(), ref_blocks.data(), nblocks * sizeof(block_q4_0)) != 0) { + for (int b = 0; b < nblocks; b++) { + const uint8_t *got = (const uint8_t *)encoded.data() + static_cast(b) * 18; + if (memcmp(got, &ref_blocks[b], 18) != 0) { + printf("block %d differs:\n got: ", b); + for (int i = 0; i < 18; i++) { + printf("%02x ", got[i]); + } + printf("\n expected: "); + for (int i = 0; i < 18; i++) { + printf("%02x ", ((const uint8_t *)&ref_blocks[b])[i]); + } + printf("\n"); + break; + } + } + printf("Encoded blocks do not match the reference\n"); + return 1; + } + + for (int i = 0; i < N; i++) { + // Compare bit patterns: the match must be exact. + uint32_t got_bits = 0, want_bits = 0; + memcpy(&got_bits, &dequantized.data()[i], sizeof(got_bits)); + memcpy(&want_bits, &ref_dequantized[i], sizeof(want_bits)); + if (got_bits != want_bits) { + printf("Dequantized values do not match the reference\n"); + return 1; + } + } + + float ref_dot = 0.0f; + for (int i = 0; i < N; i++) { + const float product = ref_dequantized[i] * acts_data[i]; + ref_dot += product; + } + printf("dot = %f (reference %f)\n", dot_result(), ref_dot); + // The dot product isn't bit-exact, since Halide may fuse the multiply + // and add differently than our C++ compiler does. Only the bytes of the + // format are contractual. + if (std::fabs(dot_result() - ref_dot) > 1e-5f * std::fabs(ref_dot)) { + printf("The dot product does not match the reference\n"); + return 1; + } + + // The fp32 scale before it was rounded to fp16 is also available, since + // we kept a handle on the quantizer: it's the `d` of the reference. + Buffer scales = fp32_scale.realize({nblocks}); + for (int b = 0; b < nblocks; b++) { + float amax = 0.0f, max = 0.0f; + for (int j = 0; j < QK4_0; j++) { + float v = weights_data[b * QK4_0 + j]; + if (amax < std::fabs(v)) { + amax = std::fabs(v); + max = v; + } + } + require(scales(b) == max / -8, "fp32 scale"); + } + + // Likewise for the values that the decoder saw: what the quantizer's + // decode direction produced, as (within, block). + Func decoded_blocks = approx.decoded_by(q.quantize); + Buffer blocks = decoded_blocks.realize({QK4_0, nblocks}); + for (int b = 0; b < nblocks; b++) { + for (int j = 0; j < QK4_0; j++) { + require(blocks(j, b) == ref_dequantized[b * QK4_0 + j], "decoded_by"); + } + } + + // ------------------------------------------------------------------------ + // Part 6: Checking what the scheme claims. + // + // Every unit declares the type and range of the values on its ports. The + // ranges are never checked when running, but we can check statically that + // a stage's guarantees imply the next one's preconditions... + // ------------------------------------------------------------------------ + require(check_ranges(q4_0).empty(), "static range check"); + + Distribution normal = Distribution::normal(0, 1); + + // ...and, on random data, get a report of how much the round trip loses. + std::cout << verify_round_trip(q4_0, normal, Float(32), {N}, 1); + + // ...and check them dynamically by running the scheme on random inputs. + // ".at(stage)" restricts a property to the values that arrive at just + // that stage, as computed by the stages before it: packing nibbles is + // lossless, provided that the quantizer really does produce [0, 15]. + require(check_property(q4_0, lossless().at(q.pack), normal, Float(32), {N}).passed, "property"); + require(check_property(q4_0, lossless().at(q.offset), normal, Float(32), {N}).passed, "property"); + + // The quantizer is lossy, but it declares a bound of one step. + require(check_property(q4_0, within_declared_bound().at(q.quantize), normal, Float(32), {N}).passed, "property"); + + printf("Success!\n"); + return 0; +}