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 @@ -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×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).
*/
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.
Expand All @@ -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<String> = 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<String> {
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<String> =
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<String> =
perLayerKv(spec) + "token"
Expand Down Expand Up @@ -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},
Expand All @@ -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()
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", "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)
Expand All @@ -53,6 +53,11 @@ public fun main(args: Array<String>) {
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")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -569,6 +578,79 @@ 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, 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`.
*/
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<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 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))
// 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)) }

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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(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)
}

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 @@ -253,6 +253,24 @@ public class GemmaModel<T : DType, V>(
public val sinSliding: Tensor<T, V>,
)

/**
* 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, 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`.
*/
public class ChunkContext<T : DType, V>(
public val cosGlobal: Tensor<T, V>,
public val sinGlobal: Tensor<T, V>,
public val cosSliding: Tensor<T, V>,
public val sinSliding: Tensor<T, V>,
public val maskGlobal: Tensor<T, V>,
public val maskSliding: Tensor<T, V>,
)

private class GemmaBlockRefs<T : DType, V>(
val attnNorm: RMSNormalization<T, V>,
val mha: MultiHeadAttention<T, V>,
Expand Down Expand Up @@ -320,6 +338,89 @@ public class GemmaModel<T : DType, V>(
/** 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<T, V>,
chunk: ChunkContext<T, V>,
selectAt: Tensor<T, V>,
selfKIn: List<Tensor<T, V>>,
selfVIn: List<Tensor<T, V>>,
ctx: ExecutionContext,
): GemmaWithPastOutput<T, V> {
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<Tensor<T, V>>(blocks.size)
val nsv = ArrayList<Tensor<T, V>>(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, nHeads, C, past+C]` (no built-in causal path). */
private fun attnWithPastChunk(
mha: MultiHeadAttention<T, V>,
sn: Tensor<T, V>,
cos: Tensor<T, V>, sin: Tensor<T, V>,
mask: Tensor<T, V>,
pastK: Tensor<T, V>, pastV: Tensor<T, V>,
ctx: ExecutionContext,
): Triple<Tensor<T, V>, Tensor<T, V>, Tensor<T, V>> {
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<T, V>, heads: Int): Tensor<T, V> =
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<T, V>,
sn: Tensor<T, V>,
Expand Down