Repository navigation
iree-android: stateful prefill/step contract with a reusable prompt-prefix KV snapshot (and an embeddings-input variant for Vulkan) #410
Description
Activity
Measured the existing graphs standalone on the arm32 Android device (2026-09-03) to pin down the contract before any runtime work. Three things the KV session must get right, plus the numbers:
1. The with-past graph has no sliding-window mask — the caller must window the cache. With the K/V from
gemma_prefill_atat SEQ 1024 for an 852-token prompt (first token48, confirmed against the all-positions redecode graph at position P−1),gemma_with_pastfed the full 852-position cache on every layer returns the wrong next token (host llvm-cpu236761; truth from the all-positions graph with the token appended:6639). Passing only the last 512 positions for the 15 sliding layers (l % 6 != 5) and the full cache for the 3 global layers gives6639on the host and on Mali. The August board loop never exceeded 13 tokens, soGemmaKvDecoder(which never trims) has the same bug for P > 512. The KV session should keep a 512-position ring per sliding layer.2. Per-step cost with a resident cache, P = 852 (windowed):
target token per step Mali (Vulkan, valhall4, embedding gathered on the host, iree-benchmark-module-vk)6639 ✓ 2.10 s (10 ms CPU) arm32 llvm-cpu, 4 threads228414 ✗ 3.31 s arm32 llvm-cpu, 4 threads, 13-position cache— 2.92 s arm32 llvm-cpu, 1 thread— 8.95 s The cost is a fixed per-token floor, not cache-bound. On the GPU it is dispatch overhead: the with-past module has 687 dispatch sites (≈ 38 per layer, ≈ 3 ms each).
--iree-dispatch-creation-enable-aggressive-fusion=truedoes not compile for SPIR-V on 3.11.0 (failed to legalize unresolved materialization … workgroup memref), and--iree-input-demote-f32-to-f16still fails on the parameter globals, so within 3.11.0 the lever is the number of dispatches the export emits per layer, not a flag.3. arm32
llvm-cpuparity failure on the with-past graph. With byte-identical inputs (windowed cache, same.irpa), host x64 and Mali agree on6639; the arm32 build returns228414(4 threads) — and171025/174631before windowing where host gave236761. The all-positions redecode graph andgemma_atare byte-identical to the host on the same device, so this is specific to the dynamic-past graph on the arm32 target. Needs a layer-by-layer dump to isolate.Also relevant: the IREE 3.12.0 nightly (
iree-3.12.0rc20260903) fixes thevector.stepfailure on the embedding gather but fails one step later on the same dispatch (failed to legalize operation 'memref.alloc' : memref<1xindex, #spirv.storage_class<Workgroup>>), so the embeddings input stays necessary for the Vulkan tier for now. Modules built by the candidate need thehalVM module at version 7; the 3.11.0 runtime provides 6.Graphs, inputs and scripts for all of the above: a local harness (happy to share).
Follow-up with the lever found after the numbers above: bf16 weight globals (the harness's default
rewriteGlobalsToBf16— bf16 parameters,stablehlo.convertto f32 on load, f32 compute; 536 MB archive instead of 1,075 MB).Same device, same windowed 852-position cache, same host-gathered embedding:
graph target FP32 archive bf16 archive token gemma_with_past, one decode tokenMali (Vulkan) 2,097 ms 257 ms 6639 ✓ both gemma_with_past, one decode tokenarm32 llvm-cpu, 4 threads 3,310 ms, wrong token (228414) 845 ms, 6639 ✓ gemma_prefill_at, SEQ 1024, 852-token promptMali — 24.8 s 48 ✓ (= all-positions graph) gemma_prefill_at, SEQ 64, 13-token promptMali 3.16 s 1.43 s 32691 ✓ So the FP32 per-token floor was the 1 GB parameter set itself (placement on Mali's unified memory; on the 32-bit CPU it even produced wrong tokens), not launch overhead — the dispatch count barely changes (655 → 591).
--iree-dispatch-creation-enable-fuse-horizontal-contractionshad no effect (655 dispatches, 2,044 ms). The arm32 "parity failure" in my previous comment is therefore withdrawn as a codegen bug: it does not reproduce with bf16 archives.Consequences for the contract:
- bf16 archives should be the shipped default for the KV graphs (they already are the harness default; FP32 is the bring-up mode).
- With the 512-ring for sliding layers, 0.26 s/token decode and 24.8 s one-time catalog prefill, a 15-token tool call projects to ≈ 13 s with the utterance fed token by token, and ≈ 5.5 s with a chunk prefill-with-past graph (
tokens 1×C + past K/V → K/V, token, C = 64) — that graph is the remaining export item for an interactive demo on this device.
Closing the loop with an in-process measurement. The three pieces are now on branches/PRs:
- functiongemma: position-selected graphs gemma_at / gemma_prefill_at — LM head on one position (#406) #415 — position-selected graphs (
gemma_at,gemma_prefill_at) - functiongemma: chunk prefill-with-past graph gemma_prefill_with_past — an utterance against the cache in one call (#410) #417 — chunk prefill-with-past graph (
gemma_prefill_with_past, C = 64 default; masks per head, see the PR for why) - iree-android: stateful KV session (IreeKvSession / IreeKvDecoder, libskainet_iree_kv.so) — prefill once, snapshot, chunk per utterance, step per token (#410) #418 — the Android KV session (
IreeKvSession/IreeKvDecoder,libskainet_iree_kv.so): three IREE sessions, device-resident K/V, zero-copy 512-position tail views for the sliding layers, native RoPE tables + chunk masks, embedding rows read from the archive, snapshot/restore without copies
Measured in-process on the arm32 Android device (Mali via Vulkan, bf16 archives, chunk 32), golden-8 German voice-command utterances with the 843-token shared tool-catalog prefix, 16 decode tokens:
stage measured session open (3 archives) 12.4 s; RSS 1.85 GB → 1.47 GB after the prefill session is released catalog prefix, 843 tokens, once per process 25.2 s per utterance: restore + one chunk call + 16 tokens 0 ms + ≈ 2.0 s + ≈ 3.9 s (245 ms/token) per utterance, total p50 5.87 s, max 6.25 s All 8 utterances return a complete tool call (or the model's refusal prose) within budget; 2/8 token streams are identical to the eager JVM run for all 16 tokens, the rest diverge after the function name (bf16 vs FP32 eager). The "embeddings input" item of this issue is what #418 implements natively (rows from the archive), so the Kotlin API stays token ids only; the arm32 CPU path runs the same code at ≈ 0.85 s/token and ≈ 5 s per 32-token chunk.
Still open from this issue: per-SEQ parameter archives (#406, second half) and the shared-archive question; neither gates the numbers above.
- functiongemma: position-selected graphs gemma_at / gemma_prefill_at — LM head on one position (#406) #415 — position-selected graphs (
Motivation, with numbers
runtime-iree-androidonly offers the stateless redecode contract (step(IntArray) -> IntArray, whole sequence recomputed). On an arm32 Android device (32-bit process, 4× ARMv8, Mali) with FunctionGemma-270M FP32 and IREE 3.11.0:A 15-token tool call against an 853-token catalog prompt therefore costs ≈ 19 min on the CPU and ≈ 7.5 min on the GPU. Prompt shortening does not fix it (short descriptions: −29 % tokens, −2/8 accuracy). With a KV contract that prefills the catalog once per process, snapshots the KV state, and per turn prefills only the utterance (≈ 36 tokens) plus ~15 decode steps, the same call projects to 8–29 s on the GPU. There is no other route to an interactive NLU on this class of device, and no arm64 firmware option (the product stays 32-bit).
Proposal
prefill(ids) -> (lastId, kvHandle),snapshot(kvHandle) -> kvHandle',step(kvHandle, id) -> id,release(kvHandle); KV kept device-side (HAL buffers), never round-tripped through the JVM. The FunctionGemma export already emitsgemma_prefill/gemma_with_pastgraphs (FunctionGemmaExportHarness.kt:258–383) — this is the runtime side of those.prefill(ids, embeddings)): IREE 3.11.0's SPIR-V backend cannot legalize the token-embedding gather (failed to legalize operation 'vector.step'on the gather dispatch). Gathering the 640-float rows on the host and passingtensor<1xSEQx640xf32>makes the same graph compile and run on Mali at 2.4–2.6× the CPU speed. The JNI currently accepts onlyIntArray.I can provide the device, the measured graphs (MLIR with the sliced head and the host-gather input), the
iree-benchmark-modulescripts, and the golden sets for parity.Context. Measured on 2026-09-03 while bringing FunctionGemma-270M up as an NLU cartridge on an arm32 Android device (Android 14,
armeabi-v7a-only 32-bit process, 4× ARMv8 @ 2.0 GHz, Mali GPU) with SKaiNET 0.53.0 + SKaiNET-transformers 0.53.0 from Maven Central. Harness and raw result files: a local harness (can share on request).