From bf9b2998672769ab900fe17950d49206a385cc34 Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Fri, 4 Sep 2026 07:54:04 +0200 Subject: [PATCH 1/2] =?UTF-8?q?functiongemma:=20chunk=20prefill-with-past?= =?UTF-8?q?=20graph=20gemma=5Fprefill=5Fwith=5Fpast=20=E2=80=94=20a=20fixe?= =?UTF-8?q?d=2064-token=20chunk=20against=20the=20dynamic=20cache=20in=20o?= =?UTF-8?q?ne=20call=20(#410)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- .../functiongemma/FunctionGemmaContract.kt | 39 ++++++- .../functiongemma/FunctionGemmaExportCli.kt | 9 +- .../FunctionGemmaExportHarness.kt | 79 ++++++++++++++ .../FunctionGemmaExportDumpTest.kt | 15 +++ .../sk/ainet/models/gemma/GemmaModel.kt | 100 ++++++++++++++++++ 5 files changed, 239 insertions(+), 3 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 cfe99804..d82dd2e1 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 @@ -41,6 +41,15 @@ public object FunctionGemmaContract { public const val FN_REDECODE_AT: String = "gemma_at" public const val FN_PREFILL_AT: String = "gemma_prefill_at" + /** + * Chunk prefill against the cache: `gemma_prefill_with_past(tokens C i32, per-base cos/sin [C×headDim], + * per-type additive masks [1×1×C×?], select [1×C], per-layer K/V …) → per-layer K/V extended by C, token 1xi32`. + * One call per utterance instead of C `gemma_with_past` steps. Fixed chunk size [DEFAULT_CHUNK]; masks + * carry causal band, padding and the sliding window (0 = attend, -1e30 = masked). + */ + public const val FN_PREFILL_WITH_PAST: String = "gemma_prefill_with_past" + public const val DEFAULT_CHUNK: Int = 64 + /** * The IREE runtime addresses functions by their module-qualified name (`module.gemma`); * `iree-run-module --function=gemma` qualifies internally, `IreeRedecodeSession` does not. @@ -56,6 +65,31 @@ public object FunctionGemmaContract { /** `gemma_at` / `gemma_prefill_at` argument order: tokens, then the one-hot position row. */ public fun selectArgs(): List = listOf("tokens", "select") + /** + * `gemma_prefill_with_past` argument order (first-use order of the trace): tokens, then per layer + * the cos/sin of its RoPE base on first use, that layer's K then V, and the layer type's mask on + * first use (it is consumed after the K/V concat); the one-hot select comes LAST. + */ + public fun prefillWithPastArgs(spec: FunctionGemmaSpec): List { + val args = mutableListOf("tokens") + var introSliding = false + var introGlobal = false + for (l in 0 until spec.nLayers) { + val isGlobal = l % spec.globalLayerPeriod == spec.globalLayerPeriod - 1 + if (isGlobal && !introGlobal) { args += "cosGlobal"; args += "sinGlobal" } + if (!isGlobal && !introSliding) { args += "cosSliding"; args += "sinSliding" } + args += "l$l.k"; args += "l$l.v" + if (isGlobal && !introGlobal) { args += "maskGlobal"; introGlobal = true } + if (!isGlobal && !introSliding) { args += "maskSliding"; introSliding = true } + } + args += "select" + return args + } + + /** `gemma_prefill_with_past` result order: per-layer K,V (extended by the chunk) then the token LAST. */ + public fun prefillWithPastOutputs(spec: FunctionGemmaSpec): List = + perLayerKv(spec) + "token" + /** `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" @@ -100,7 +134,7 @@ public object FunctionGemmaContract { |{ | "contractVersion": $CONTRACT_VERSION, | "model": "functiongemma-270m", - | "functions": { "redecode": "$FN_REDECODE", "prefill": "$FN_PREFILL", "withPast": "$FN_WITH_PAST", "redecodeAt": "$FN_REDECODE_AT", "prefillAt": "$FN_PREFILL_AT" }, + | "functions": { "redecode": "$FN_REDECODE", "prefill": "$FN_PREFILL", "withPast": "$FN_WITH_PAST", "redecodeAt": "$FN_REDECODE_AT", "prefillAt": "$FN_PREFILL_AT", "prefillWithPast": "$FN_PREFILL_WITH_PAST" }, | "nLayers": ${spec.nLayers}, | "headDim": ${spec.headDim}, | "nKvHeads": ${spec.nKvHeads}, @@ -122,6 +156,9 @@ public object FunctionGemmaContract { | "redecodeAtArgs": [${arr(selectArgs())}], | "prefillAtArgs": [${arr(selectArgs())}], | "prefillAtOutputs": [${arr(prefillAtOutputs(spec))}], + | "chunk": $DEFAULT_CHUNK, + | "prefillWithPastArgs": [${arr(prefillWithPastArgs(spec))}], + | "prefillWithPastOutputs": [${arr(prefillWithPastOutputs(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 900f2bda..7abef5dd 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", "redecode_at", "prefill_at", "all")) { - error("GEMMA_GRAPH must be redecode|prefill|with_past|redecode_at|prefill_at|all, got '$graph'") + if (graph !in setOf("redecode", "prefill", "with_past", "redecode_at", "prefill_at", "prefill_with_past", "all")) { + error("GEMMA_GRAPH must be redecode|prefill|with_past|redecode_at|prefill_at|prefill_with_past|all, got '$graph'") } val spec = FunctionGemmaSpec(gguf = gguf, seq = seq, partialRotary = partial, quant = quant) @@ -53,6 +53,11 @@ public fun main(args: Array) { 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_with_past") { + val chunk = System.getenv("GEMMA_CHUNK")?.toInt() ?: FunctionGemmaContract.DEFAULT_CHUNK + FunctionGemmaExportHarness.exportPrefillWithPast(spec, outDir, chunk) + println("[functiongemma-export] prefill_with_past: wrote $outDir/gemma-prefill-with-past.mlir + gemma-prefill-with-past.safetensors (chunk=$chunk)") + } if (graph == "prefill_at") { FunctionGemmaExportHarness.exportPrefillAt(spec, outDir) println("[functiongemma-export] prefill_at: wrote $outDir/gemma-prefill-at.mlir + gemma-prefill-at.safetensors") 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 33c30a5a..e6ed5fa5 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 @@ -143,6 +143,15 @@ public object FunctionGemmaExportHarness { bf16 = spec.quant != FunctionGemmaQuant.FP32, ) + /** The chunk prefill-with-past graph (`func @gemma_prefill_with_past`, chunk = [FunctionGemmaContract.DEFAULT_CHUNK]) from [spec]. */ + public fun exportPrefillWithPast(spec: FunctionGemmaSpec, outDir: String, chunk: Int = FunctionGemmaContract.DEFAULT_CHUNK): String = exportPrefillWithPast( + gguf = spec.gguf, + outDir = outDir, + chunk = chunk, + 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() @@ -569,6 +578,76 @@ public object FunctionGemmaExportHarness { mlir } + /** + * KV-cache CHUNK PREFILL graph `func @gemma_prefill_with_past`: a fixed [chunk] of token ids against the + * dynamic per-layer past cache ([GemmaModel.forwardPrefillWithPast]). Inputs beyond the cache are + * host-built and position-carrying: per-base cos/sin `[chunk, headDim]`, per-layer-type additive + * masks `[1, 1, chunk, past+chunk]` (dynamic last dim), and the one-hot `select [1, chunk]` for the + * last real token. Results: per-layer K/V extended by the chunk (`1x{nKV}x?x{headDim}`, padding + * positions included — slice like after prefill), then `token 1xi32`. Writes + * `gemma-prefill-with-past.mlir` + its own `gemma-prefill-with-past.safetensors`. + */ + public fun exportPrefillWithPast( + gguf: String, + outDir: String, + chunk: Int = FunctionGemmaContract.DEFAULT_CHUNK, + 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 nLayers = md.blockCount + val headDim = md.getHeadDim(0) + val nKV = md.kvHeadCount + val sentinelRollback = System.getenv("GEMMA_SENTINEL_PAST") == "1" + val pastDim = if (sentinelRollback) SENTINEL_PAST else sk.ainet.lang.tensor.Dim.DYNAMIC + val kvDim = if (sentinelRollback) SENTINEL_PAST + chunk else sk.ainet.lang.tensor.Dim.DYNAMIC + + val tokens = voidF32(Shape(chunk)) + val cosG = voidF32(Shape(chunk, headDim)); val sinG = voidF32(Shape(chunk, headDim)) + val cosS = voidF32(Shape(chunk, headDim)); val sinS = voidF32(Shape(chunk, headDim)) + val maskG = voidF32(Shape(1, 1, chunk, kvDim)); val maskS = voidF32(Shape(1, 1, chunk, kvDim)) + val select = voidF32(Shape(1, chunk)) + val selfKIn = List(nLayers) { voidF32(Shape(1, nKV, pastDim, headDim)) } + val selfVIn = List(nLayers) { voidF32(Shape(1, nKV, pastDim, headDim)) } + + 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.forwardPrefillWithPast( + tokens, GemmaModel.ChunkContext(cosG, sinG, cosS, sinS, maskG, maskS), select, selfKIn, selfVIn, ectx, + ) + ectx.ops.argMax(out.logits, dim = -1) // [1] i32 — the selected position's token + } 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_WITH_PAST) + var mlir = if (bf16) rewriteGlobalsToBf16(module.content) else module.content + if (sentinelRollback) mlir = relaxSeqDimToDynamic(mlir) + File(outDir).apply { mkdirs() } + File(outDir, "gemma-prefill-with-past.mlir").writeText(mlir) + writeSafetensors(module.externalParameters, File(outDir, "gemma-prefill-with-past.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 86bcb97f..dd72c969 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 @@ -76,6 +76,21 @@ class FunctionGemmaExportDumpTest { checkResultCount(prefillAtMlir, FunctionGemmaContract.FN_PREFILL_AT, FunctionGemmaContract.prefillAtOutputs(spec).size) } + @Test + fun chunkPrefillWithPast_emitsContractShapes() { + FunctionGemmaFixture.assumeRealCheckpointRunnable() + val outDir = File(System.getProperty("java.io.tmpdir"), "functiongemma-export-dump-chunk").absolutePath + FunctionGemmaExportHarness.exportPrefillWithPast(spec, outDir) + val mlir = File(outDir, "gemma-prefill-with-past.mlir").readText() + val c = FunctionGemmaContract.DEFAULT_CHUNK + assertTrue(mlir.contains("func.func @${FunctionGemmaContract.FN_PREFILL_WITH_PAST}("), "entry func present") + assertTrue(mlir.contains("x?x${spec.headDim}"), "dynamic past cache dim must be present") + assertTrue(mlir.contains("tensor<1x1x${c}x?xf32>"), "additive mask inputs [1,1,C,?] present") + assertTrue(mlir.contains("tensor<${c}x${spec.headDim}xf32>"), "per-position cos/sin tables present") + checkArgCount(mlir, FunctionGemmaContract.FN_PREFILL_WITH_PAST, FunctionGemmaContract.prefillWithPastArgs(spec).size) + checkResultCount(mlir, FunctionGemmaContract.FN_PREFILL_WITH_PAST, FunctionGemmaContract.prefillWithPastOutputs(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 41c1777b..ea11ce22 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 @@ -253,6 +253,23 @@ public class GemmaModel( public val sinSliding: Tensor, ) + /** + * Inputs of [forwardPrefillWithPast] for a C-token chunk: per-RoPE-base cos/sin tables `[C, headDim]` + * (one row per absolute position `past .. past+C-1`, split-half layout as in [buildRopeCosSin]) and + * per-layer-type **additive** attention masks `[1, 1, C, past+C]` (0 = attend, -1e30 = masked). The + * masks carry everything position-dependent — the causal band inside the chunk, the zero padding + * beyond the real tokens, and the 512-position sliding window of the sliding layers — so the graph + * itself stays position-agnostic like `gemma_with_past`. + */ + public class ChunkContext( + public val cosGlobal: Tensor, + public val sinGlobal: Tensor, + public val cosSliding: Tensor, + public val sinSliding: Tensor, + public val maskGlobal: Tensor, + public val maskSliding: Tensor, + ) + private class GemmaBlockRefs( val attnNorm: RMSNormalization, val mha: MultiHeadAttention, @@ -320,6 +337,89 @@ public class GemmaModel( /** Hand-wired single-new-token self-attention over the past cache for one block. RoPE applies * at the runtime position via the fed-in [cos]/[sin]. Returns (attnOut `[1, qDim]`, * extendedK `[1, nKVHeads, past+1, headDim]`, extendedV). */ + /** + * KV-cache CHUNK PREFILL: run a fixed-size chunk [tokens] (`[C]` token ids, zero-padded) against the + * incoming per-layer self cache (`[1, nKVHeads, past, headDim]`) → logits `[1, vocab]` for the ONE + * position picked by the one-hot [selectAt] (`[1, C]`), plus the per-layer cache extended by all C + * positions (`[1, nKVHeads, past+C, headDim]`; the caller slices the padding off, as after prefill). + * This is what makes a user utterance cost one call instead of C single-token [forwardWithPast] steps. + */ + public fun forwardPrefillWithPast( + tokens: Tensor, + chunk: ChunkContext, + selectAt: Tensor, + selfKIn: List>, + selfVIn: List>, + ctx: ExecutionContext, + ): GemmaWithPastOutput { + require(ple == null) { "GemmaModel.forwardPrefillWithPast: PLE models not supported by the KV-cache path yet" } + val ops = ctx.ops + val rawEmbeds = tokenEmbedding.forward(tokens, ctx) // [C, hidden] + var h = if (embedScale != 1f) ops.mulScalar(rawEmbeds, embedScale) else rawEmbeds + val nsk = ArrayList>(blocks.size) + val nsv = ArrayList>(blocks.size) + for ((i, block) in blocks.withIndex()) { + val r = refsFor(block) + val isGlobal = r.mha.slidingWindow == null + val cos = if (isGlobal) chunk.cosGlobal else chunk.cosSliding + val sin = if (isGlobal) chunk.sinGlobal else chunk.sinSliding + val mask = if (isGlobal) chunk.maskGlobal else chunk.maskSliding + val sn = r.attnNorm.forward(h, ctx) + val (attnOut, fullK, fullV) = attnWithPastChunk(r.mha, sn, cos, sin, mask, selfKIn[i], selfVIn[i], ctx) + val postAttn = r.postAttnNorm?.forward(attnOut, ctx) ?: attnOut + 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 + nsk += ops.reshape(fullK, fullK.shape) // identity → distinct graph output node + nsv += ops.reshape(fullV, fullV.shape) + } + val normed = outputNorm.forward(h, ctx) // [C, hidden] + val hAt = ops.matmul(selectAt, normed) // [1, hidden] + val logits = applySoftcap(lmHead.forward(hAt, ctx), ctx) // [1, vocab] + return GemmaWithPastOutput(logits, nsk, nsv) + } + + /** [attnWithPast] for a C-row chunk: heads-first `[heads, C, headDim]` projections, RoPE from `[C, headDim]` + * tables, and the caller's additive mask `[1, 1, C, past+C]` (no built-in causal path). */ + private fun attnWithPastChunk( + mha: MultiHeadAttention, + sn: Tensor, + cos: Tensor, sin: Tensor, + mask: Tensor, + pastK: Tensor, pastV: Tensor, + ctx: ExecutionContext, + ): Triple, Tensor, Tensor> { + val ops = ctx.ops + val rope = mha.rope ?: error("GemmaModel.attnWithPastChunk: MHA has no RoPE") + val idx = paramIdx(mha) + val c = sn.shape[0] + fun headsFirst(x: Tensor, heads: Int): Tensor = + ops.permute(ops.reshape(x, Shape(c, heads, mha.headDim)), intArrayOf(1, 0, 2)) // [heads, C, headDim] + var q = headsFirst(linearProject(ops, sn, mha.params[idx[0]].value), mha.nHeads) + var k = headsFirst(linearProject(ops, sn, mha.params[idx[1]].value), mha.nKVHeads) + val v = headsFirst(linearProject(ops, sn, mha.params[idx[2]].value), mha.nKVHeads) + val qn = mha.qNorm; val kn = mha.kNorm + if (qn != null && kn != null) { q = qn.forward(q, ctx); k = kn.forward(k, ctx) } // qkNorm BEFORE RoPE + val cos3 = ops.unsqueeze(cos, 0); val sin3 = ops.unsqueeze(sin, 0) // [1, C, headDim] over heads + q = rope.forwardWithCosSin(q, cos3, sin3, ctx) + k = rope.forwardWithCosSin(k, cos3, sin3, ctx) + val fullK = ops.concat(listOf(pastK, ops.unsqueeze(k, 0)), dim = 2) // [1, nKVHeads, past+C, headDim] + val fullV = ops.concat(listOf(pastV, ops.unsqueeze(v, 0)), dim = 2) + val eK = expandKV(fullK, mha.nHeads, mha.nKVHeads, ops) + val eV = expandKV(fullV, mha.nHeads, mha.nKVHeads, ops) + val scale = mha.attentionScale ?: (1f / sqrt(mha.headDim.toFloat())) + val o = ops.scaledDotProductAttention( + query = ops.unsqueeze(q, 0), key = eK, value = eV, + mask = mask, scale = scale, causal = false, // the additive mask carries causal + padding + window + ) // [1, nHeads, C, headDim] + val merged = ops.reshape(ops.permute(ops.squeeze(o, 0), intArrayOf(1, 0, 2)), Shape(c, mha.nHeads * mha.headDim)) + val attnOut = linearProject(ops, merged, mha.params[idx[3]].value) // o_proj → [C, hidden] + return Triple(attnOut, fullK, fullV) + } + private fun attnWithPast( mha: MultiHeadAttention, sn: Tensor, From ee1c4d32faedd6033d38257eac1094544bd09fcc Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Fri, 4 Sep 2026 07:58:38 +0200 Subject: [PATCH 2/2] =?UTF-8?q?functiongemma:=20chunk=20masks=20are=20per-?= =?UTF-8?q?head=20[1,=20nHeads,=20C,=20past+C]=20=E2=80=94=20a=20broadcast?= =?UTF-8?q?=20over=20heads=20to=20a=20dynamic=20shape=20is=20not=20express?= =?UTF-8?q?ible=20in=20static=20StableHLO?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../sk/ainet/models/functiongemma/FunctionGemmaContract.kt | 2 +- .../models/functiongemma/FunctionGemmaExportHarness.kt | 7 +++++-- .../models/functiongemma/FunctionGemmaExportDumpTest.kt | 2 +- .../commonMain/kotlin/sk/ainet/models/gemma/GemmaModel.kt | 5 +++-- 4 files changed, 10 insertions(+), 6 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 d82dd2e1..80a24c9d 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 @@ -43,7 +43,7 @@ public object FunctionGemmaContract { /** * Chunk prefill against the cache: `gemma_prefill_with_past(tokens C i32, per-base cos/sin [C×headDim], - * per-type additive masks [1×1×C×?], select [1×C], per-layer K/V …) → per-layer K/V extended by C, token 1xi32`. + * per-type additive masks [1×nHeads×C×?], select [1×C], per-layer K/V …) → per-layer K/V extended by C, token 1xi32`. * One call per utterance instead of C `gemma_with_past` steps. Fixed chunk size [DEFAULT_CHUNK]; masks * carry causal band, padding and the sliding window (0 = attend, -1e30 = masked). */ 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 e6ed5fa5..753c7301 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 @@ -582,7 +582,7 @@ public object FunctionGemmaExportHarness { * KV-cache CHUNK PREFILL graph `func @gemma_prefill_with_past`: a fixed [chunk] of token ids against the * dynamic per-layer past cache ([GemmaModel.forwardPrefillWithPast]). Inputs beyond the cache are * host-built and position-carrying: per-base cos/sin `[chunk, headDim]`, per-layer-type additive - * masks `[1, 1, chunk, past+chunk]` (dynamic last dim), and the one-hot `select [1, chunk]` for the + * masks `[1, nHeads, chunk, past+chunk]` (dynamic last dim; per head, see the note in the body), and the one-hot `select [1, chunk]` for the * last real token. Results: per-layer K/V extended by the chunk (`1x{nKV}x?x{headDim}`, padding * positions included — slice like after prefill), then `token 1xi32`. Writes * `gemma-prefill-with-past.mlir` + its own `gemma-prefill-with-past.safetensors`. @@ -616,7 +616,10 @@ public object FunctionGemmaExportHarness { val tokens = voidF32(Shape(chunk)) val cosG = voidF32(Shape(chunk, headDim)); val sinG = voidF32(Shape(chunk, headDim)) val cosS = voidF32(Shape(chunk, headDim)); val sinS = voidF32(Shape(chunk, headDim)) - val maskG = voidF32(Shape(1, 1, chunk, kvDim)); val maskS = voidF32(Shape(1, 1, chunk, kvDim)) + // Masks are per-head [1, nHeads, C, past+C]: a [1, 1, C, ?] mask would need a broadcast to a + // dynamic shape inside the attention add, which StableHLO's static broadcast_in_dim cannot express. + val nHeads = md.headCount + val maskG = voidF32(Shape(1, nHeads, chunk, kvDim)); val maskS = voidF32(Shape(1, nHeads, chunk, kvDim)) val select = voidF32(Shape(1, chunk)) val selfKIn = List(nLayers) { voidF32(Shape(1, nKV, pastDim, headDim)) } val selfVIn = List(nLayers) { voidF32(Shape(1, nKV, pastDim, headDim)) } 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 dd72c969..9fd64242 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 @@ -85,7 +85,7 @@ class FunctionGemmaExportDumpTest { val c = FunctionGemmaContract.DEFAULT_CHUNK assertTrue(mlir.contains("func.func @${FunctionGemmaContract.FN_PREFILL_WITH_PAST}("), "entry func present") assertTrue(mlir.contains("x?x${spec.headDim}"), "dynamic past cache dim must be present") - assertTrue(mlir.contains("tensor<1x1x${c}x?xf32>"), "additive mask inputs [1,1,C,?] present") + assertTrue(Regex("tensor<1x\\d+x${c}x\\?xf32>").containsMatchIn(mlir), "additive mask inputs [1,nHeads,C,?] present") assertTrue(mlir.contains("tensor<${c}x${spec.headDim}xf32>"), "per-position cos/sin tables present") checkArgCount(mlir, FunctionGemmaContract.FN_PREFILL_WITH_PAST, FunctionGemmaContract.prefillWithPastArgs(spec).size) checkResultCount(mlir, FunctionGemmaContract.FN_PREFILL_WITH_PAST, FunctionGemmaContract.prefillWithPastOutputs(spec).size) 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 ea11ce22..81e33bfa 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 @@ -256,7 +256,8 @@ public class GemmaModel( /** * Inputs of [forwardPrefillWithPast] for a C-token chunk: per-RoPE-base cos/sin tables `[C, headDim]` * (one row per absolute position `past .. past+C-1`, split-half layout as in [buildRopeCosSin]) and - * per-layer-type **additive** attention masks `[1, 1, C, past+C]` (0 = attend, -1e30 = masked). The + * per-layer-type **additive** attention masks `[1, nHeads, C, past+C]` (0 = attend, -1e30 = masked; + * per head because a broadcast over heads to a dynamic shape is not expressible in static StableHLO). The * masks carry everything position-dependent — the causal band inside the chunk, the zero padding * beyond the real tokens, and the 512-position sliding window of the sliding layers — so the graph * itself stays position-agnostic like `gemma_with_past`. @@ -383,7 +384,7 @@ public class GemmaModel( } /** [attnWithPast] for a C-row chunk: heads-first `[heads, C, headDim]` projections, RoPE from `[C, headDim]` - * tables, and the caller's additive mask `[1, 1, C, past+C]` (no built-in causal path). */ + * tables, and the caller's additive mask `[1, nHeads, C, past+C]` (no built-in causal path). */ private fun attnWithPastChunk( mha: MultiHeadAttention, sn: Tensor,