Skip to content

iree-android: stateful prefill/step contract with a reusable prompt-prefix KV snapshot (and an embeddings-input variant for Vulkan) #410

Description

@michalharakal

Motivation, with numbers

runtime-iree-android only 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:

SEQ CPU arm32 per step Vulkan per step (head sliced, embeddings gathered on the host)
64 10.4 s (6.6 s with the head sliced) 3.16 s
256 41.9 s 7.76 s
1024 77.3 s (sliced head) 29.9 s

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

  1. Contract: 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 emits gemma_prefill/gemma_with_past graphs (FunctionGemmaExportHarness.kt:258–383) — this is the runtime side of those.
  2. Head sliced to the target position in every graph (separate issue; 36 % of the step and the 32-bit memory cap).
  3. Embeddings-input variant (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 passing tensor<1xSEQx640xf32> makes the same graph compile and run on Mali at 2.4–2.6× the CPU speed. The JNI currently accepts only IntArray.
  4. Error reporting through JNI (see the null-return issue).

I can provide the device, the measured graphs (MLIR with the sliced head and the host-gather input), the iree-benchmark-module scripts, 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).

Activity

  1. michalharakal commented on Sep 3, 2026

    @michalharakal
    ContributorAuthor

    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_at at SEQ 1024 for an 852-token prompt (first token 48, confirmed against the all-positions redecode graph at position P−1), gemma_with_past fed the full 852-position cache on every layer returns the wrong next token (host llvm-cpu 236761; 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 gives 6639 on the host and on Mali. The August board loop never exceeded 13 tokens, so GemmaKvDecoder (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 threads 228414 ✗ 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=true does not compile for SPIR-V on 3.11.0 (failed to legalize unresolved materialization … workgroup memref), and --iree-input-demote-f32-to-f16 still 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-cpu parity failure on the with-past graph. With byte-identical inputs (windowed cache, same .irpa), host x64 and Mali agree on 6639; the arm32 build returns 228414 (4 threads) — and 171025 / 174631 before windowing where host gave 236761. The all-positions redecode graph and gemma_at are 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 the vector.step failure 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 the hal VM 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).

  2. michalharakal commented on Sep 3, 2026

    @michalharakal
    ContributorAuthor

    Follow-up with the lever found after the numbers above: bf16 weight globals (the harness's default rewriteGlobalsToBf16 — bf16 parameters, stablehlo.convert to 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 token Mali (Vulkan) 2,097 ms 257 ms 6639 ✓ both
    gemma_with_past, one decode token arm32 llvm-cpu, 4 threads 3,310 ms, wrong token (228414) 845 ms, 6639 ✓
    gemma_prefill_at, SEQ 1024, 852-token prompt Mali — 24.8 s 48 ✓ (= all-positions graph)
    gemma_prefill_at, SEQ 64, 13-token prompt Mali 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-contractions had 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.
  3. michalharakal commented on Sep 4, 2026

    @michalharakal
    ContributorAuthor

    Closing the loop with an in-process measurement. The three pieces are now on branches/PRs:

    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.

  4. added 2 commits that reference this issue on Sep 4, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions