From 95f2d278c2254719f8a7d261a1b7a3a4969fa226 Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Thu, 3 Sep 2026 22:47:07 +0200 Subject: [PATCH] =?UTF-8?q?functiongemma:=20position-selected=20graphs=20g?= =?UTF-8?q?emma=5Fat=20/=20gemma=5Fprefill=5Fat=20=E2=80=94=20LM=20head=20?= =?UTF-8?q?on=20one=20one-hot-selected=20position,=20single-token=20result?= =?UTF-8?q?=20(#406)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 --- .../functiongemma/FunctionGemmaContract.kt | 27 +++- .../functiongemma/FunctionGemmaExportCli.kt | 12 +- .../FunctionGemmaExportHarness.kt | 135 ++++++++++++++++++ .../FunctionGemmaExportDumpTest.kt | 23 +++ .../sk/ainet/models/gemma/GemmaModel.kt | 53 ++++++- 5 files changed, 246 insertions(+), 4 deletions(-) diff --git a/llm-inference/functiongemma/src/commonMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaContract.kt b/llm-inference/functiongemma/src/commonMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaContract.kt index c203346d..cfe99804 100644 --- a/llm-inference/functiongemma/src/commonMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaContract.kt +++ b/llm-inference/functiongemma/src/commonMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaContract.kt @@ -32,12 +32,34 @@ public object FunctionGemmaContract { public const val FN_PREFILL: String = "gemma_prefill" public const val FN_WITH_PAST: String = "gemma_with_past" + /** + * Position-selected variants (contract addendum, 2026-09): same trunk, LM head applied to ONE + * position chosen by a one-hot `select [1, seq]` f32 input, single `token 1xi32` result. They + * remove the `seq x vocab` logits/argmax scratch of the all-positions graphs (a 1024-position + * graph otherwise needs ~1 GB of scratch and cannot run in a 32-bit process). + */ + public const val FN_REDECODE_AT: String = "gemma_at" + public const val FN_PREFILL_AT: String = "gemma_prefill_at" + + /** + * The IREE runtime addresses functions by their module-qualified name (`module.gemma`); + * `iree-run-module --function=gemma` qualifies internally, `IreeRedecodeSession` does not. + */ + public fun qualified(fn: String): String = if ('.' in fn) fn else "module.$fn" + /** Both graphs emit K THEN V per block (board-verified; see GemmaKvDecoder.kFirstInOutput). */ public const val K_FIRST_IN_OUTPUT: Boolean = true /** `gemma_prefill` / `gemma` argument order. */ public fun prefillArgs(): List = listOf("tokens") + /** `gemma_at` / `gemma_prefill_at` argument order: tokens, then the one-hot position row. */ + public fun selectArgs(): List = listOf("tokens", "select") + + /** `gemma_prefill_at` result order: per-layer K,V caches then the single selected token LAST. */ + public fun prefillAtOutputs(spec: FunctionGemmaSpec): List = + perLayerKv(spec) + "token" + /** `gemma_prefill` result order: per-layer K,V caches then the argMax tokens LAST. */ public fun prefillOutputs(spec: FunctionGemmaSpec): List = perLayerKv(spec) + "tokens" @@ -78,7 +100,7 @@ public object FunctionGemmaContract { |{ | "contractVersion": $CONTRACT_VERSION, | "model": "functiongemma-270m", - | "functions": { "redecode": "$FN_REDECODE", "prefill": "$FN_PREFILL", "withPast": "$FN_WITH_PAST" }, + | "functions": { "redecode": "$FN_REDECODE", "prefill": "$FN_PREFILL", "withPast": "$FN_WITH_PAST", "redecodeAt": "$FN_REDECODE_AT", "prefillAt": "$FN_PREFILL_AT" }, | "nLayers": ${spec.nLayers}, | "headDim": ${spec.headDim}, | "nKvHeads": ${spec.nKvHeads}, @@ -97,6 +119,9 @@ public object FunctionGemmaContract { | "prefillOutputs": [${arr(prefillOutputs(spec))}], | "withPastArgs": [${arr(withPastArgs(spec))}], | "withPastOutputs": [${arr(withPastOutputs(spec))}], + | "redecodeAtArgs": [${arr(selectArgs())}], + | "prefillAtArgs": [${arr(selectArgs())}], + | "prefillAtOutputs": [${arr(prefillAtOutputs(spec))}], | "toolMap": { $tools } |} |""".trimMargin() diff --git a/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportCli.kt b/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportCli.kt index 9c6f56f5..900f2bda 100644 --- a/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportCli.kt +++ b/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportCli.kt @@ -32,8 +32,8 @@ public fun main(args: Array) { else -> FunctionGemmaQuant.BF16 } val graph = (System.getenv("GEMMA_GRAPH") ?: "redecode").lowercase() - if (graph !in setOf("redecode", "prefill", "with_past", "all")) { - error("GEMMA_GRAPH must be redecode|prefill|with_past|all, got '$graph'") + if (graph !in setOf("redecode", "prefill", "with_past", "redecode_at", "prefill_at", "all")) { + error("GEMMA_GRAPH must be redecode|prefill|with_past|redecode_at|prefill_at|all, got '$graph'") } val spec = FunctionGemmaSpec(gguf = gguf, seq = seq, partialRotary = partial, quant = quant) @@ -49,6 +49,14 @@ public fun main(args: Array) { FunctionGemmaExportHarness.exportPrefill(spec, outDir) println("[functiongemma-export] prefill: wrote $outDir/gemma-prefill.mlir (seq=$seq)") } + if (graph == "redecode_at") { + val r = FunctionGemmaExportHarness.exportRedecodeAt(spec, outDir) + println("[functiongemma-export] redecode_at: wrote ${r.mlirPath} + ${r.safetensorsPath} (${r.externalParamCount} externals, ${r.weightMiB} MiB, seq=${r.seq})") + } + if (graph == "prefill_at") { + FunctionGemmaExportHarness.exportPrefillAt(spec, outDir) + println("[functiongemma-export] prefill_at: wrote $outDir/gemma-prefill-at.mlir + gemma-prefill-at.safetensors") + } if (graph == "with_past" || graph == "all") { FunctionGemmaExportHarness.exportWithPast(spec, outDir) println("[functiongemma-export] with_past: wrote $outDir/gemma-with-past.mlir (dynamic 1x1x?x256 cache)") diff --git a/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportHarness.kt b/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportHarness.kt index e48731c4..33c30a5a 100644 --- a/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportHarness.kt +++ b/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportHarness.kt @@ -125,6 +125,24 @@ public object FunctionGemmaExportHarness { bf16 = spec.quant != FunctionGemmaQuant.FP32, ) + /** The position-selected re-decode graph (`func @gemma_at`) from [spec]. */ + public fun exportRedecodeAt(spec: FunctionGemmaSpec, outDir: String): RedecodeResult = exportRedecodeAt( + gguf = spec.gguf, + outDir = outDir, + seq = spec.seq, + partialRotary = spec.partialRotary, + bf16 = spec.quant != FunctionGemmaQuant.FP32, + ) + + /** The position-selected prefill graph (`func @gemma_prefill_at`) from [spec]. */ + public fun exportPrefillAt(spec: FunctionGemmaSpec, outDir: String): String = exportPrefillAt( + gguf = spec.gguf, + outDir = outDir, + seq = spec.seq, + partialRotary = spec.partialRotary, + bf16 = spec.quant != FunctionGemmaQuant.FP32, + ) + /** Write `manifest.json` (see [FunctionGemmaContract.manifestJson]). */ public fun writeManifest(spec: FunctionGemmaSpec, outDir: String): File { File(outDir).mkdirs() @@ -434,6 +452,123 @@ public object FunctionGemmaExportHarness { mlir } + /** + * Position-selected re-decode graph `func @gemma_at(tokens 1x{seq} i32, select 1x{seq} f32) -> token 1xi32`: + * the same fixed-[seq] pass as [export], but the LM head runs on ONE position picked by the one-hot + * `select` row ([GemmaModel.forwardAt]). Measured on a 4-core arm32 box at seq 64: 6.6 s/step vs + * 10.4 s for the all-positions graph, and the 1024-position graph fits a 32-bit process (the + * all-positions variant dies on a 1 GB `seq x vocab` argmax scratch). f32 or bf16 externals only. + */ + public fun exportRedecodeAt( + gguf: String, + outDir: String, + seq: Int = 24, + partialRotary: Float = 1.0f, + bf16: Boolean = true, + ): RedecodeResult = runBlocking { + val ctx = DirectCpuExecutionContext.create() + val weights = GemmaWeightLoader( + randomAccessProvider = { JvmRandomAccessSource.open(gguf) }, + weightForm = GEMMA_DEQUANTIZE_ALL, + ).loadToMapStreaming(ctx, FP32::class) + val patched = weights.copy( + metadata = weights.metadata.copy( + ropeParametersFull = weights.metadata.ropeParametersFull.copy(partialRotaryFactor = partialRotary), + ), + ) + @Suppress("UNCHECKED_CAST") + val model = GemmaNetworkLoader.fromWeights(ctx, patched, FP32::class) as GemmaModel + fun stripKvCache(m: Module<*, *>) { + if (m is MultiHeadAttention<*, *>) m.kvCache = null + m.modules.forEach { stripKvCache(it) } + } + stripKvCache(model) + + val input = voidF32(Shape(1, seq)) // token ids -> `1x{seq}xi32` + val select = voidF32(Shape(1, seq)) // one-hot position row -> `1x{seq}xf32` + val tapeCtx = DefaultGraphExecutionContext.tape(baseOps = VoidTensorOps()) + val tape = tapeCtx.record { + val ct = (this as DefaultGraphExecutionContext).currentTape ?: error("no tape") + Execution.tapeStack.pushTape(ct) + try { + val ectx = this as ExecutionContext + val logits = model.forwardAt(input, select, ectx) // [1, vocab] f32 + ectx.ops.argMax(logits, dim = -1) // [1] i32 + } finally { + Execution.tapeStack.popTape() + } + }.first + val graph = (tape as DefaultExecutionTape).toComputeGraph(synthesizeExternalInputs = true, embedConstants = true) + val module = StableHloConverterFactory + .createBasic(ConstantMaterializationPolicy.ExternalAlways(scope = "model")) + .convert(graph, FunctionGemmaContract.FN_REDECODE_AT) + val out = File(outDir).apply { mkdirs() } + val ext = module.externalParameters + val mlir = if (bf16) rewriteGlobalsToBf16(module.content) else module.content + val mlirFile = File(out, "gemma-gen-at.mlir").apply { writeText(mlir) } + val stFile = File(out, "gemma-at.safetensors") + writeSafetensors(ext, stFile, bf16) + val totalF32 = ext.sumOf { it.source.sizeInBytes } + RedecodeResult( + mlirPath = mlirFile.absolutePath, + safetensorsPath = stFile.absolutePath, + externalParamCount = ext.size, + weightMiB = (if (bf16) totalF32 / 2 else totalF32) / (1024 * 1024), + seq = seq, + ) + } + + /** + * Position-selected PREFILL graph `func @gemma_prefill_at(tokens {seq} i32, select 1x{seq} f32)` -> + * per-layer initial K/V (`1x{nKV}x{seq}x{headDim}`) THEN `token 1xi32` last. Same K/V contract as + * [exportPrefill]; the caller passes the one-hot row for position `P-1` instead of reading it out + * of a `{seq}xi32` argmax. Writes `gemma-prefill-at.mlir` + its own `gemma-prefill-at.safetensors`. + */ + public fun exportPrefillAt( + gguf: String, + outDir: String, + seq: Int = 24, + partialRotary: Float = 1.0f, + bf16: Boolean = true, + ): String = runBlocking { + val ctx = DirectCpuExecutionContext.create() + val weights = GemmaWeightLoader( + randomAccessProvider = { JvmRandomAccessSource.open(gguf) }, + weightForm = GEMMA_DEQUANTIZE_ALL, + ).loadToMapStreaming(ctx, FP32::class) + val md = weights.metadata + val patched = weights.copy( + metadata = md.copy(ropeParametersFull = md.ropeParametersFull.copy(partialRotaryFactor = partialRotary)), + ) + @Suppress("UNCHECKED_CAST") + val model = GemmaNetworkLoader.fromWeights(ctx, patched, FP32::class) as GemmaModel + + val tokens = voidF32(Shape(seq)) // [seq] token ids -> `{seq}xi32` graph input + val select = voidF32(Shape(1, seq)) // one-hot position row -> `1x{seq}xf32` + val tapeCtx = DefaultGraphExecutionContext.tape(baseOps = VoidTensorOps()) + val tape = tapeCtx.record { + val ct = (this as DefaultGraphExecutionContext).currentTape ?: error("no tape") + Execution.tapeStack.pushTape(ct) + try { + val ectx = this as ExecutionContext + val out = model.forwardPrefillAt(tokens, select, ectx) + ectx.ops.argMax(out.logits, dim = -1) // [1] i32 — the selected position's token + // out.selfK / out.selfV are terminal -> per-layer initial K/V outputs. + } finally { + Execution.tapeStack.popTape() + } + }.first + val graph = (tape as DefaultExecutionTape).toComputeGraph(synthesizeExternalInputs = true, embedConstants = true) + val module = StableHloConverterFactory + .createBasic(ConstantMaterializationPolicy.ExternalAlways(scope = "model")) + .convert(graph, FunctionGemmaContract.FN_PREFILL_AT) + val mlir = if (bf16) rewriteGlobalsToBf16(module.content) else module.content + File(outDir).apply { mkdirs() } + File(outDir, "gemma-prefill-at.mlir").writeText(mlir) + writeSafetensors(module.externalParameters, File(outDir, "gemma-prefill-at.safetensors"), bf16) + mlir + } + /** * Write [ext] as a safetensors archive ([bf16] truncation = core parity, else raw f32). Every traced * graph numbers its "model" externals independently (`t0`, `t10`, …), so an archive only serves the diff --git a/llm-inference/functiongemma/src/jvmTest/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportDumpTest.kt b/llm-inference/functiongemma/src/jvmTest/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportDumpTest.kt index 408878db..86bcb97f 100644 --- a/llm-inference/functiongemma/src/jvmTest/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportDumpTest.kt +++ b/llm-inference/functiongemma/src/jvmTest/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportDumpTest.kt @@ -53,6 +53,29 @@ class FunctionGemmaExportDumpTest { "withPast=${File(outDir, "gemma-with-past.safetensors").length() / (1 shl 20)}MiB manifest=${r.manifest}") } + @Test + fun positionSelectedGraphs_emitContractShapes() { + FunctionGemmaFixture.assumeRealCheckpointRunnable() + val outDir = File(System.getProperty("java.io.tmpdir"), "functiongemma-export-dump-at").absolutePath + + // ---- gemma_at: (tokens 1xseq i32, select 1xseq f32) -> token 1xi32, fully static ---- + val r = FunctionGemmaExportHarness.exportRedecodeAt(spec, outDir) + val atMlir = File(r.mlirPath).readText() + assertTrue(!atMlir.contains("tensor"), "one-hot select input present") + assertTrue(!atMlir.contains("x${spec.seq}x262144xf32>") && !atMlir.contains("x${spec.seq}x262144xi32>"), "gemma_at: no seq x vocab logits/argmax scratch") + checkArgCount(atMlir, FunctionGemmaContract.FN_REDECODE_AT, FunctionGemmaContract.selectArgs().size) + checkResultCount(atMlir, FunctionGemmaContract.FN_REDECODE_AT, 1) + + // ---- gemma_prefill_at: (tokens seq i32, select 1xseq f32) -> per-layer K/V THEN token last ---- + val prefillAtMlir = File(outDir, "gemma-prefill-at.mlir").also { FunctionGemmaExportHarness.exportPrefillAt(spec, outDir) }.readText() + assertTrue(!prefillAtMlir.contains("tensor"), "gemma_prefill_at: no seq x vocab logits") + checkArgCount(prefillAtMlir, FunctionGemmaContract.FN_PREFILL_AT, FunctionGemmaContract.selectArgs().size) + checkResultCount(prefillAtMlir, FunctionGemmaContract.FN_PREFILL_AT, FunctionGemmaContract.prefillAtOutputs(spec).size) + } + private fun checkArgCount(mlir: String, func: String, expected: Int) { val sig = signatureLine(mlir, func) val argsPart = sig.substringAfter("@$func(").substringBefore(") ->") diff --git a/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaModel.kt b/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaModel.kt index fb5c073d..41c1777b 100644 --- a/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaModel.kt +++ b/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaModel.kt @@ -78,7 +78,20 @@ public class GemmaModel( add(lmHead) } - override fun onForward(input: Tensor, ctx: ExecutionContext): Tensor { + override fun onForward(input: Tensor, ctx: ExecutionContext): Tensor = forwardImpl(input, ctx, null) + + /** + * [forward] with the LM head applied to ONE position: [selectAt] is a one-hot `[1, seq]` row + * that multiplies the final-normed hidden state `[seq, hidden]` down to `[1, hidden]` before + * `lm_head`, so the graph returns logits `[1, vocab]` instead of `[seq, vocab]`. For a compiled + * re-decode graph this removes the `seq × vocab` logits and argmax scratch that only one position + * ever reads (36 % of a FunctionGemma step on a 4-core arm32 box; the reason a 1024-position + * graph cannot run in a 32-bit process). Plain matmul on purpose: no dynamic-index op is needed. + */ + public fun forwardAt(input: Tensor, selectAt: Tensor, ctx: ExecutionContext): Tensor = + forwardImpl(input, ctx, selectAt) + + private fun forwardImpl(input: Tensor, ctx: ExecutionContext, selectAt: Tensor?): Tensor { // Diagnostic: dump per-block hidden state stats when GEMMA4_DUMP_HIDDEN=1. // Compares against expected llama.cpp magnitudes to localize the // BOS-loop forward-pass bug. See gemma4-research/findings/dsl_vs_llamacpp_logit_divergence.md. @@ -125,6 +138,11 @@ public class GemmaModel( // Step 4: final norm + lm_head. hidden = outputNorm.forward(hidden, ctx) if (dumpHidden) dumpHiddenStats("post-norm ", hidden) + if (selectAt != null) { + // [1, seq] x [seq, hidden] -> [1, hidden]; a rank-3 [1, seq, hidden] trunk output is squeezed first. + val h2 = if (hidden.rank == 3) ctx.ops.squeeze(hidden, dim = 0) else hidden + hidden = ctx.ops.matmul(selectAt, h2) + } var logits = lmHead.forward(hidden, ctx) if (dumpHidden) dumpHiddenStats("logits-pre-sc", logits) @@ -379,6 +397,39 @@ public class GemmaModel( return GemmaPrefillOutput(logits, selfK, selfV) } + /** + * [forwardPrefill] with the LM head applied to ONE position ([selectAt] = one-hot `[1, seq]`, see + * [forwardAt]): logits `[1, vocab]` + the same per-layer initial self K/V. The compiled + * `gemma_prefill_at` graph reads the first generated token from this single row instead of + * computing `seq × vocab` logits for the padding positions. + */ + public fun forwardPrefillAt(input: Tensor, selectAt: Tensor, ctx: ExecutionContext): GemmaPrefillOutput { + require(ple == null) { "GemmaModel.forwardPrefillAt: PLE models not supported by the KV-cache path yet" } + stripCaches() + val ops = ctx.ops + val rawEmbeds = tokenEmbedding.forward(input, ctx) + var h = if (embedScale != 1f) ops.mulScalar(rawEmbeds, embedScale) else rawEmbeds + val selfK = ArrayList>(blocks.size) + val selfV = ArrayList>(blocks.size) + for (block in blocks) { + val r = refsFor(block) + val kv = r.mha.forwardWithKV(r.attnNorm.forward(h, ctx), null, ctx) + val postAttn = r.postAttnNorm?.forward(kv.output, ctx) ?: kv.output + val h1 = ops.add(h, postAttn) + val ffnOut = r.ffn.forward(r.ffnNorm.forward(h1, ctx), ctx) + val postFfw = r.postFfwNorm?.forward(ffnOut, ctx) ?: ffnOut + var h2 = ops.add(h1, postFfw) + if (r.outScale != null) h2 = r.outScale.forward(h2, ctx) + h = h2 + selfK += ops.unsqueeze(kv.k, 0) + selfV += ops.unsqueeze(kv.v, 0) + } + val normed = outputNorm.forward(h, ctx) + val hAt = ops.matmul(selectAt, if (normed.rank == 3) ops.squeeze(normed, dim = 0) else normed) // [1, hidden] + val logits = applySoftcap(lmHead.forward(hAt, ctx), ctx) // [1, vocab] + return GemmaPrefillOutput(logits, selfK, selfV) + } + /** * KV-cache DECODE step: one new [tokenId] (`[1]` or `[1,1]`) at the position encoded by [rope] * ([buildRopeCosSin]); the incoming per-layer self cache ([selfKIn]/[selfVIn],