Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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<String> = listOf("tokens")

/** `gemma_at` / `gemma_prefill_at` argument order: tokens, then the one-hot position row. */
public fun selectArgs(): List<String> = 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<String> =
perLayerKv(spec) + "token"

/** `gemma_prefill` result order: per-layer K,V caches then the argMax tokens LAST. */
public fun prefillOutputs(spec: FunctionGemmaSpec): List<String> =
perLayerKv(spec) + "tokens"
Expand Down Expand Up @@ -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},
Expand All @@ -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()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,8 +32,8 @@ public fun main(args: Array<String>) {
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)
Expand All @@ -49,6 +49,14 @@ public fun main(args: Array<String>) {
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)")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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<FP32, Float>(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<FP32, Float>
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<FP32, Float>(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<FP32, Float>

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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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<?"), "gemma_at: dynamic shape leaked into the fixed graph")
assertTrue(atMlir.contains("func.func @${FunctionGemmaContract.FN_REDECODE_AT}("), "entry func present")
assertTrue(atMlir.contains("tensor<1x${spec.seq}xf32>"), "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: must stay fully static (fixed seq)")
assertTrue(!prefillAtMlir.contains("x${spec.seq}x262144xf32>"), "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(") ->")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -78,7 +78,20 @@ public class GemmaModel<T : DType, V>(
add(lmHead)
}

override fun onForward(input: Tensor<T, V>, ctx: ExecutionContext): Tensor<T, V> {
override fun onForward(input: Tensor<T, V>, ctx: ExecutionContext): Tensor<T, V> = 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<T, V>, selectAt: Tensor<T, V>, ctx: ExecutionContext): Tensor<T, V> =
forwardImpl(input, ctx, selectAt)

private fun forwardImpl(input: Tensor<T, V>, ctx: ExecutionContext, selectAt: Tensor<T, V>?): Tensor<T, V> {
// 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.
Expand Down Expand Up @@ -125,6 +138,11 @@ public class GemmaModel<T : DType, V>(
// 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)

Expand Down Expand Up @@ -379,6 +397,39 @@ public class GemmaModel<T : DType, V>(
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<T, V>, selectAt: Tensor<T, V>, ctx: ExecutionContext): GemmaPrefillOutput<T, V> {
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<Tensor<T, V>>(blocks.size)
val selfV = ArrayList<Tensor<T, V>>(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],
Expand Down