Skip to content

functiongemma: chunk prefill-with-past graph gemma_prefill_with_past — an utterance against the cache in one call (#410) - #417

Merged
michalharakal merged 2 commits into
developfrom
functiongemma-chunk-prefill
Sep 4, 2026
Merged

michalharakal merged 2 commits into
developfrom
functiongemma-chunk-prefill

Conversation

@michalharakal

@michalharakal michalharakal commented Sep 4, 2026 •

Copy link
Copy Markdown
Contributor

Part of #410 (the "chunk prefill" item). Stacked on #415 — merge that first.

Why

With the KV contract, a user utterance still costs one gemma_with_past call per token (0.26 s each on the device's Mali GPU with bf16 archives → 9–14 s for a 36–52-token utterance). This graph runs a fixed-size chunk of tokens against the cache in one call.

Change (additive)

  • GemmaModel.forwardPrefillWithPast(tokens, ChunkContext, selectAt, selfKIn, selfVIn, ctx) and attnWithPastChunk: heads-first [heads, C, headDim] projections (reshape + permute), RoPE from per-position [C, headDim] cos/sin tables, attention over past + C keys with a caller-supplied additive mask per layer type [1, nHeads, C, past+C] (0 = attend, -1e30 = masked; the mask carries the causal band inside the chunk, the zero padding, and the 512-position sliding window, so the graph stays position-agnostic like gemma_with_past), K/V extended by the chunk, LM head on the one-hot-selected position.
  • Masks are per head on purpose: a [1, 1, C, ?] mask needs a broadcast over heads to a dynamic shape inside the attention add, which static stablehlo.broadcast_in_dim cannot express (the first export failed there).
  • FunctionGemmaExportHarness.exportPrefillWithPast(gguf, outDir, chunk = 64, …) → gemma-prefill-with-past.{mlir,safetensors}; spec overload; CLI GEMMA_GRAPH=prefill_with_past (GEMMA_CHUNK).
  • Contract: FN_PREFILL_WITH_PAST, DEFAULT_CHUNK = 64, prefillWithPastArgs(spec) (first-use order: tokens, then per layer the cos/sin of its base on first use, K, V, and the type's mask on first use; select last — verified against the emitted signature: 44 args, 37 results for the 18-layer spec), prefillWithPastOutputs, manifest entries. Dump test.

Validation (device: arm32 Android device, IREE 3.11.0, bf16 archive)

The 852-token tool-catalog prompt split as an 800-position cache (from gemma_prefill_at) plus a 52-token chunk (C = 64, padded), select on row 51 — must reproduce the token the all-positions redecode graph gives at position 851 (48):

target token one chunk call
host llvm-cpu x64 48 ✓ —
Mali (Vulkan, valhall4, embeddings gathered on the host) 48 ✓ 4.58 s
arm32 llvm-cpu, 4 threads 48 ✓ 9.28 s

Same check with GEMMA_CHUNK=32 (820 cached positions + a 32-token chunk): token 48 ✓ on all three targets, Mali 2.45 s, arm32 5.12 s — the chunk cost scales with C, so the chunk size is a deployment choice driven by the utterance-length distribution.

With 0.26 s per decode token, a 15-token tool call with a warm catalog cache projects to ≈ 8.5 s (C = 64) or ≈ 6.4 s (C = 32) on the GPU.

…— a fixed 64-token chunk against the dynamic cache in one call (#410)

GemmaModel.forwardPrefillWithPast / attnWithPastChunk: heads-first [heads, C, headDim] projections,
RoPE from per-position [C, headDim] tables, attention over past+C with a caller-built additive
mask [1, 1, C, past+C] per layer type (causal band, padding, and the 512 sliding window all live
in the mask, so the graph stays position-agnostic), K/V extended by the chunk, LM head on the
one-hot-selected position. Contract: FN_PREFILL_WITH_PAST, DEFAULT_CHUNK = 64,
prefillWithPastArgs/Outputs (first-use order), manifest entries; CLI GEMMA_GRAPH=prefill_with_past
(GEMMA_CHUNK); dump test.
…oadcast over heads to a dynamic shape is not expressible in static StableHLO
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant