Skip to content

functiongemma: position-selected graphs gemma_at / gemma_prefill_at — LM head on one position (#406) - #415

Merged
michalharakal merged 1 commit into
developfrom
functiongemma-position-selected-head
Sep 4, 2026
Merged

michalharakal merged 1 commit into
developfrom
functiongemma-position-selected-head

Conversation

@michalharakal

@michalharakal michalharakal commented Sep 3, 2026 •

Copy link
Copy Markdown
Contributor

Addresses #406 (the LM-head half). Stacked on #414 — merge that first.

Why

The exported redecode and prefill graphs run lm_head and the argmax over all SEQ positions and then materialise three more SEQ×262144 tensors for the argmax, although the caller reads one position per step. Measured on an arm32 Android device (arm32, 4 threads, FP32): that is 36 % of every step, and at SEQ 1024 the 1024×262144 i32 scratch (1.07 GB) cannot be allocated in a 32-bit process — the catalog-sized prompt could not run at all.

Change (additive; the existing graphs are untouched)

  • GemmaModel.forwardAt(input, selectAt, ctx) / forwardPrefillAt(...): a one-hot [1, seq] row multiplies the final-normed hidden state down to [1, hidden] before lm_head — a plain matmul, so no dynamic-index op is needed and every backend lowers it.
  • FunctionGemmaExportHarness.exportRedecodeAt/exportPrefillAt: gemma_at(tokens 1×SEQ i32, select 1×SEQ f32) → 1×i32 and gemma_prefill_at(tokens SEQ i32, select) → per-layer K/V…, 1×i32 (own archives, like the other graphs).
  • Contract: FN_REDECODE_AT, FN_PREFILL_AT, selectArgs(), prefillAtOutputs(), qualified() (the runtime wants module.<fn>; see iree-android: nativeStep returns null silently on any failure; FunctionGemmaContract.FN_REDECODE ("gemma") is not the qualified name the runtime needs ("module.gemma") #404), manifest entries redecodeAt/prefillAt + arg/result arrays. CONTRACT_VERSION unchanged (nothing existing moves).
  • CLI: GEMMA_GRAPH=redecode_at|prefill_at.
  • FunctionGemmaExportDumpTest.positionSelectedGraphs_emitContractShapes.

Measured

Not in this PR

The per-SEQ parameter archives (the other half of #406) and the sliding-window contract for the with-past graph (#410).

… LM head on one one-hot-selected position, single-token result (#406)

GemmaModel.forwardAt / forwardPrefillAt multiply the final-normed hidden state by a one-hot
[1, seq] row before lm_head (plain matmul, no dynamic-index op), so the compiled graphs return
[1, vocab] logits and a 1xi32 token instead of seq x vocab logits plus a seq x vocab argmax
scratch. Measured on a 4-core arm32 box at seq 64: 6.6 s/step vs 10.4 s; the 1024-position
graph now fits a 32-bit process (the all-positions variant died on a 1 GB calloc).
Contract additions are additive: FN_REDECODE_AT, FN_PREFILL_AT, selectArgs(),
prefillAtOutputs(), manifest entries, qualified(); CLI GEMMA_GRAPH=redecode_at|prefill_at;
dump test for both graphs.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
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