diff --git a/llm-apps/skainet-cli/src/main/kotlin/sk/ainet/apps/skainet/cli/Main.kt b/llm-apps/skainet-cli/src/main/kotlin/sk/ainet/apps/skainet/cli/Main.kt index 858251e4..2212fba3 100644 --- a/llm-apps/skainet-cli/src/main/kotlin/sk/ainet/apps/skainet/cli/Main.kt +++ b/llm-apps/skainet-cli/src/main/kotlin/sk/ainet/apps/skainet/cli/Main.kt @@ -33,6 +33,10 @@ import sk.ainet.exec.kernel.NativeTernaryF32GemvKernel import sk.ainet.exec.kernel.NativeTernaryLmheadKernel import sk.ainet.models.bitnet.BitNetNetworkLoader import sk.ainet.models.bitnet.BitNetPackedGgufLoader +import sk.ainet.models.bitnet.BitNetTwoStageDecode +import sk.ainet.models.bitnet.bitnetPlanesHead +import sk.ainet.models.bitnet.generateTwoStage +import sk.ainet.lang.tensor.data.BitNetPlanesTensorData import sk.ainet.models.qwen.QwenNetworkLoader import java.nio.file.Path import kotlin.io.path.exists @@ -230,6 +234,10 @@ fun main(args: Array) { // activation, QK-Norm, or ungated FFN, so logits silently // diverged from the checkpoint's intent. The DSL path is correct // for Apertus too. See APERTUS_ROLLOUT.md (PR 1). + // Set on the BitNet path when the lm_head loaded as BITNET_PLANES — gates the + // two-stage decode loop at the generation site (transformers#358). + var bitnetTwoStageHead: BitNetPlanesTensorData? = null + val runtime: InferenceRuntime = if (modelInfo.family == ModelFamily.GEMMA) { println("Loading Gemma GGUF model from $modelPath via gemmaNetwork() + OptimizedLLMRuntime (engine loader, keep-packed, mapped)...") if (cliArgs.contextLength != null) { @@ -266,6 +274,13 @@ fun main(args: Array) { ctx, sourceProvider = { JvmRandomAccessSource.open(modelPath.toString()) }, ) + bitnetTwoStageHead = bitnetPlanesHead(loaded.model) + if (bitnetTwoStageHead != null) { + println( + "Two-stage lm_head decode: enabled (BITNET_PLANES head, " + + "top-${BitNetTwoStageDecode.DEFAULT_CANDIDATES} exact rescoring).", + ) + } OptimizedLLMRuntime( model = loaded.model, ctx = ctx, @@ -363,9 +378,30 @@ fun main(args: Array) { println("---") print(promptText) + val twoStageHead = bitnetTwoStageHead + @Suppress("UNCHECKED_CAST") + val twoStageRuntime = if (twoStageHead != null) runtime as? OptimizedLLMRuntime else null val elapsed = measureTime { - runtime.generate(prompt = promptTokens, steps = cliArgs.steps, temperature = cliArgs.temperature) { id -> - print(tokenizer.decode(id)) + if (twoStageHead != null && twoStageRuntime != null) { + // transformers#358: trunk-only forwards + fused stage-1 scan + exact top-200 + // rescoring — the full-vocab lm_head matmul never runs, prefill included. + twoStageRuntime.generateTwoStage( + prompt = promptTokens, + steps = cliArgs.steps, + temperature = cliArgs.temperature, + head = twoStageHead, + native = if (NativeTernaryLmheadKernel.isAvailable()) { + sk.ainet.models.bitnet.BitNetStage1Kernel(NativeTernaryLmheadKernel::lmheadStage1) + } else { + null + }, + ) { id -> + print(tokenizer.decode(id)) + } + } else { + runtime.generate(prompt = promptTokens, steps = cliArgs.steps, temperature = cliArgs.temperature) { id -> + print(tokenizer.decode(id)) + } } }.inWholeMilliseconds diff --git a/llm-core/api/jvm/llm-core.api b/llm-core/api/jvm/llm-core.api index c086ad46..27218e1e 100644 --- a/llm-core/api/jvm/llm-core.api +++ b/llm-core/api/jvm/llm-core.api @@ -78,7 +78,9 @@ public final class sk/ainet/apps/llm/GenerateResult { public final class sk/ainet/apps/llm/GenerateUntilStopKt { public static final fun generateUntilStop (Lsk/ainet/apps/llm/InferenceRuntime;[IIIFLkotlin/random/Random;Lkotlin/jvm/functions/Function1;Lkotlin/jvm/functions/Function1;Lkotlin/jvm/functions/Function2;Lsk/ainet/apps/llm/PrefillStrategy;)Lsk/ainet/apps/llm/GenerateResult; + public static final fun generateUntilStop (Lsk/ainet/apps/llm/InferenceRuntime;[IILjava/util/Set;FIFLkotlin/random/Random;Lkotlin/jvm/functions/Function1;Lkotlin/jvm/functions/Function1;Lkotlin/jvm/functions/Function2;Lsk/ainet/apps/llm/PrefillStrategy;)Lsk/ainet/apps/llm/GenerateResult; public static synthetic fun generateUntilStop$default (Lsk/ainet/apps/llm/InferenceRuntime;[IIIFLkotlin/random/Random;Lkotlin/jvm/functions/Function1;Lkotlin/jvm/functions/Function1;Lkotlin/jvm/functions/Function2;Lsk/ainet/apps/llm/PrefillStrategy;ILjava/lang/Object;)Lsk/ainet/apps/llm/GenerateResult; + public static synthetic fun generateUntilStop$default (Lsk/ainet/apps/llm/InferenceRuntime;[IILjava/util/Set;FIFLkotlin/random/Random;Lkotlin/jvm/functions/Function1;Lkotlin/jvm/functions/Function1;Lkotlin/jvm/functions/Function2;Lsk/ainet/apps/llm/PrefillStrategy;ILjava/lang/Object;)Lsk/ainet/apps/llm/GenerateResult; public static final fun sampleFromLogits (Lsk/ainet/lang/tensor/Tensor;FLkotlin/random/Random;)I public static synthetic fun sampleFromLogits$default (Lsk/ainet/lang/tensor/Tensor;FLkotlin/random/Random;ILjava/lang/Object;)I } @@ -173,6 +175,7 @@ public final class sk/ainet/apps/llm/OptimizedLLMRuntime : sk/ainet/apps/llm/Inf public static synthetic fun compileWith$default (Lsk/ainet/apps/llm/OptimizedLLMRuntime;ILsk/ainet/compile/opt/GraphOptimizationPipeline;ILjava/lang/Object;)Ljava/util/List; public fun forward (I)Lsk/ainet/lang/tensor/Tensor; public fun forwardBatched ([I)Lsk/ainet/lang/tensor/Tensor; + public final fun forwardHidden (I)Lsk/ainet/lang/tensor/Tensor; public final fun generate ([IIFLkotlin/jvm/functions/Function1;)V public final fun getBos ()I public final fun getDim ()I @@ -242,7 +245,9 @@ public final class sk/ainet/apps/llm/RopeUtilsKt { public final class sk/ainet/apps/llm/SamplingUtilsKt { public static final fun sampleFromCandidates (Ljava/util/List;FLkotlin/random/Random;)I public static synthetic fun sampleFromCandidates$default (Ljava/util/List;FLkotlin/random/Random;ILjava/lang/Object;)I + public static final fun sampleFromLogits ([FFIFLkotlin/random/Random;)I public static final fun sampleFromLogits ([FFLkotlin/random/Random;)I + public static synthetic fun sampleFromLogits$default ([FFIFLkotlin/random/Random;ILjava/lang/Object;)I public static synthetic fun sampleFromLogits$default ([FFLkotlin/random/Random;ILjava/lang/Object;)I public static final fun softmaxInPlace ([FI)V public static synthetic fun softmaxInPlace$default ([FIILjava/lang/Object;)V @@ -355,6 +360,7 @@ public final class sk/ainet/apps/llm/tokenizer/GGUFTokenizer : sk/ainet/apps/llm public fun getEosTokenId ()I public final fun getTokenizerType ()Lsk/ainet/apps/llm/TokenizerType; public fun getVocabSize ()I + public final fun tokenId (Ljava/lang/String;)Ljava/lang/Integer; } public final class sk/ainet/apps/llm/tokenizer/GGUFTokenizer$Companion { diff --git a/llm-core/src/commonMain/kotlin/sk/ainet/apps/llm/OptimizedLLMRuntime.kt b/llm-core/src/commonMain/kotlin/sk/ainet/apps/llm/OptimizedLLMRuntime.kt index 49d532f7..785fde08 100644 --- a/llm-core/src/commonMain/kotlin/sk/ainet/apps/llm/OptimizedLLMRuntime.kt +++ b/llm-core/src/commonMain/kotlin/sk/ainet/apps/llm/OptimizedLLMRuntime.kt @@ -173,6 +173,45 @@ public class OptimizedLLMRuntime( return logits } + /** + * Single-token forward that stops **before the model's final module** — the lm_head + * projection — returning the last-position hidden state (`[1, dim]`) instead of logits + * (transformers#358). + * + * This is the pre-lm_head hook a candidate-based decode needs: score the vocabulary + * outside the module tree (e.g. `BitNetTwoStageDecode` over a `BITNET_PLANES` head) and + * hand `sampleFromCandidates` a candidate list, never materializing full-vocab logits. + * It is equally useful for prompt ingestion, where logits are discarded anyway — skipping + * the head saves a full-vocab projection per prefill position. + * + * State semantics are identical to [forward]: the KV cache and [position] advance exactly + * as they would on a full step — only the final, stateless projection is skipped. The top + * module executes its children in order ([sk.ainet.lang.nn.topology.MLP] semantics), so + * this runs `modules.dropLast(1)` inside the same per-step forward scope; like a logits + * tensor, the returned hidden state is only valid until the next forward begins. + * + * Module-tree path only (DIRECT / HYBRID): OPTIMIZED mode compiles the whole graph, + * head included — there is no partial execution to offer there. + */ + public fun forwardHidden(tokenId: Int): Tensor { + require(position < seqLen) { "Context length exceeded: pos=$position seqLen=$seqLen" } + require(mode != OptimizedLLMMode.OPTIMIZED) { + "forwardHidden runs the module tree; OPTIMIZED mode executes the compiled full graph" + } + val trunk = model.modules + require(trunk.size >= 2) { + "forwardHidden needs a top-level module pipeline ending in the lm_head; " + + "got ${trunk.size} child module(s)" + } + var hidden = createTokenTensor(tokenId) + val stepCtx = stepCtx() + for (i in 0 until trunk.size - 1) { + hidden = trunk[i].forward(hidden, stepCtx) + } + position++ + return hidden + } + /** * Single batched forward over [tokenIds]. Builds an `[N]`-shaped input * tensor, feeds it through the model in one shot, and returns the diff --git a/llm-inference/bitnet/api/jvm/bitnet.api b/llm-inference/bitnet/api/jvm/bitnet.api index 234a819b..5f3b0dc5 100644 --- a/llm-inference/bitnet/api/jvm/bitnet.api +++ b/llm-inference/bitnet/api/jvm/bitnet.api @@ -83,12 +83,23 @@ public final class sk/ainet/models/bitnet/BitNetPackedGgufLoader$Loaded { public fun toString ()Ljava/lang/String; } +public abstract interface class sk/ainet/models/bitnet/BitNetStage1Kernel { + public abstract fun stage1 ([FI[BIIIII[FI)V +} + public final class sk/ainet/models/bitnet/BitNetTwoStageDecode { public static final field DEFAULT_CANDIDATES I public static final field INSTANCE Lsk/ainet/models/bitnet/BitNetTwoStageDecode; public final fun exactScore (Lsk/ainet/lang/tensor/data/BitNetPlanesTensorData;[FI)F - public final fun stage1Scores (Lsk/ainet/lang/tensor/data/BitNetPlanesTensorData;[F)[F - public final fun topK (Lsk/ainet/lang/tensor/data/BitNetPlanesTensorData;[FII)Ljava/util/List; - public static synthetic fun topK$default (Lsk/ainet/models/bitnet/BitNetTwoStageDecode;Lsk/ainet/lang/tensor/data/BitNetPlanesTensorData;[FIIILjava/lang/Object;)Ljava/util/List; + public final fun stage1Scores (Lsk/ainet/lang/tensor/data/BitNetPlanesTensorData;[FLsk/ainet/models/bitnet/BitNetStage1Kernel;)[F + public static synthetic fun stage1Scores$default (Lsk/ainet/models/bitnet/BitNetTwoStageDecode;Lsk/ainet/lang/tensor/data/BitNetPlanesTensorData;[FLsk/ainet/models/bitnet/BitNetStage1Kernel;ILjava/lang/Object;)[F + public final fun topK (Lsk/ainet/lang/tensor/data/BitNetPlanesTensorData;[FIILsk/ainet/models/bitnet/BitNetStage1Kernel;)Ljava/util/List; + public static synthetic fun topK$default (Lsk/ainet/models/bitnet/BitNetTwoStageDecode;Lsk/ainet/lang/tensor/data/BitNetPlanesTensorData;[FIILsk/ainet/models/bitnet/BitNetStage1Kernel;ILjava/lang/Object;)Ljava/util/List; +} + +public final class sk/ainet/models/bitnet/BitNetTwoStageGenerationKt { + public static final fun bitnetPlanesHead (Lsk/ainet/lang/nn/Module;)Lsk/ainet/lang/tensor/data/BitNetPlanesTensorData; + public static final fun generateTwoStage (Lsk/ainet/apps/llm/OptimizedLLMRuntime;[IIFLsk/ainet/lang/tensor/data/BitNetPlanesTensorData;Lsk/ainet/models/bitnet/BitNetStage1Kernel;ILkotlin/random/Random;Lkotlin/jvm/functions/Function1;)V + public static synthetic fun generateTwoStage$default (Lsk/ainet/apps/llm/OptimizedLLMRuntime;[IIFLsk/ainet/lang/tensor/data/BitNetPlanesTensorData;Lsk/ainet/models/bitnet/BitNetStage1Kernel;ILkotlin/random/Random;Lkotlin/jvm/functions/Function1;ILjava/lang/Object;)V } diff --git a/llm-inference/bitnet/src/commonMain/kotlin/sk/ainet/models/bitnet/BitNetTwoStageDecode.kt b/llm-inference/bitnet/src/commonMain/kotlin/sk/ainet/models/bitnet/BitNetTwoStageDecode.kt index 0ef453e7..ea74ed3f 100644 --- a/llm-inference/bitnet/src/commonMain/kotlin/sk/ainet/models/bitnet/BitNetTwoStageDecode.kt +++ b/llm-inference/bitnet/src/commonMain/kotlin/sk/ainet/models/bitnet/BitNetTwoStageDecode.kt @@ -6,6 +6,28 @@ import sk.ainet.lang.memory.TernaryCodec import sk.ainet.lang.tensor.data.BitNetPlanesTensorData import sk.ainet.lang.tensor.storage.TensorEncoding +/** + * Stage-1 scorer contract: one fused pass over planes 0–3 with the FP16 row scales applied — + * the signature of the engine's `TernaryLmheadNative.lmheadStage1` symbol, restated here so + * this module's common code needs no dependency on the backend-api artifact (which does not + * publish every native target this module builds for). A JVM caller plugs the vendored kernel + * in as a method reference: `BitNetStage1Kernel(NativeTernaryLmheadKernel::lmheadStage1)`. + */ +public fun interface BitNetStage1Kernel { + public fun stage1( + activation: FloatArray, + activationOffset: Int, + weight: ByteArray, + planesByteOffset: Int, + planeStrideBytes: Int, + rowScaleByteOffset: Int, + inputDim: Int, + outputDim: Int, + out: FloatArray, + outOffset: Int, + ) +} + /** * NeoGPU's two-stage lm_head (transformers#337, upstream `hs_ml_infer.c` Stage 1/Stage 2): score * the full vocabulary cheaply with **planes 0–3** of a [TensorEncoding.BITNET_PLANES] weight, @@ -38,11 +60,32 @@ public object BitNetTwoStageDecode { * Stage-1 scores — planes 0–3 with the FP16 row scales applied — for all [weight] rows * against [hidden]. This is exactly what the fused NEON kernel computes in one pass; the * Kotlin loop here is the portable reference of the same contract. + * + * Pass [native] (e.g. the FFM `NativeTernaryLmheadKernel`, when available) to run stage 1 + * as **one** fused `lmhead_stage1` call — the same symbol the exact planes matmul makes two + * calls to, so a stage-1-only scan halves the kernel work per decode step (#358). `null` + * keeps the portable Kotlin loop; both produce the same contract. */ - public fun stage1Scores(weight: BitNetPlanesTensorData, hidden: FloatArray): FloatArray { + public fun stage1Scores( + weight: BitNetPlanesTensorData, + hidden: FloatArray, + native: BitNetStage1Kernel? = null, + ): FloatArray { val n = weight.rows val k = weight.cols require(hidden.size == k) { "hidden has ${hidden.size} elements, weight expects $k" } + if (native != null) { + val out = FloatArray(n) + native.stage1( + activation = hidden, activationOffset = 0, + weight = weight.packedData, planesByteOffset = 0, + planeStrideBytes = TensorEncoding.BITNET_PLANES.planeStrideBytes(n, k), + rowScaleByteOffset = TensorEncoding.BITNET_PLANES.rowScalesByteOffset(n, k), + inputDim = k, outputDim = n, + out = out, outOffset = 0, + ) + return out + } val bytes = weight.packedData val planeStride = TensorEncoding.BITNET_PLANES.planeStrideBytes(n, k) val rowBytes = k / 4 @@ -90,10 +133,12 @@ public object BitNetTwoStageDecode { hidden: FloatArray, k: Int = 8, maxCandidates: Int = DEFAULT_CANDIDATES, + /** Optional fused stage-1 kernel — see [stage1Scores]; stage 2 stays the exact decode. */ + native: BitNetStage1Kernel? = null, ): List { val n = weight.rows require(k in 1..n) { "k=$k outside 1..$n" } - val stage1 = stage1Scores(weight, hidden) + val stage1 = stage1Scores(weight, hidden, native) // Per-row error bound: |exact - stage1| <= rowScale * TAIL_WEIGHT * Σ|h|. var hAbs = 0f diff --git a/llm-inference/bitnet/src/commonMain/kotlin/sk/ainet/models/bitnet/BitNetTwoStageGeneration.kt b/llm-inference/bitnet/src/commonMain/kotlin/sk/ainet/models/bitnet/BitNetTwoStageGeneration.kt new file mode 100644 index 00000000..ddec15b0 --- /dev/null +++ b/llm-inference/bitnet/src/commonMain/kotlin/sk/ainet/models/bitnet/BitNetTwoStageGeneration.kt @@ -0,0 +1,97 @@ +package sk.ainet.models.bitnet + +import kotlin.random.Random +import sk.ainet.apps.llm.OptimizedLLMRuntime +import sk.ainet.apps.llm.sampleFromCandidates +import sk.ainet.lang.memory.ExperimentalMemoryApi +import sk.ainet.lang.nn.Module +import sk.ainet.lang.nn.topology.ModuleParameters +import sk.ainet.lang.tensor.data.BitNetPlanesTensorData +import sk.ainet.lang.types.FP32 + +/** + * The `BITNET_PLANES` lm_head weight of [model], or `null` when the model's final module is not + * a planes-encoded head — the gate a caller checks before reaching for [generateTwoStage]. + * + * The decoder network template ends in the lm_head projection ([bitnetNetwork] ends in the + * `"output"` `VoidDense`); its weight is planes-encoded when the model was loaded through + * [BitNetPackedGgufLoader] with `planesLmHead` on (both the `output.weight` and the tied-2B4T + * lanes, transformers#337/#357). + */ +@OptIn(ExperimentalMemoryApi::class) +public fun bitnetPlanesHead(model: Module): BitNetPlanesTensorData? { + val head = model.modules.lastOrNull() ?: return null + @Suppress("UNCHECKED_CAST") + val params = (head as? ModuleParameters)?.params ?: return null + val weight = params.firstOrNull { it.name.endsWith("weight") }?.value ?: return null + return weight.data as? BitNetPlanesTensorData +} + +/** + * The two-stage BitNet generation loop (transformers#358): NeoGPU's decode driver expressed + * over [OptimizedLLMRuntime.forwardHidden] + [BitNetTwoStageDecode] + [sampleFromCandidates]. + * + * Per step, instead of the full-vocab lm_head matmul inside the module tree: + * + * 1. [OptimizedLLMRuntime.forwardHidden] runs the trunk — everything but the head — for the + * last-position hidden state (KV cache and position advance exactly as a full forward). + * 2. [BitNetTwoStageDecode.topK] scans all rows with planes 0–3 (one fused `lmhead_stage1` + * call when [native] is given) and exactly rescores the surviving [candidates] rows. + * 3. [sampleFromCandidates] samples with the standard temperature semantics over that list. + * + * Prompt ingestion also runs [OptimizedLLMRuntime.forwardHidden]: prefill logits are discarded + * anyway, so every prompt position saves its full-vocab projection outright. + * + * Exactness: greedy decode (`temperature <= 1e-6`) selects the same token as the full matmul — + * [BitNetTwoStageDecode.topK]'s bound guarantees the exact top-k (up to the [candidates] cap, + * NeoGPU's flat-200 heuristic beyond it). At temperature, sampling is restricted to the + * rescored candidates — the two-stage contract: mass outside the top-[candidates] is treated + * as zero. + * + * Same prompt semantics as the stock `generate` loop: [OptimizedLLMRuntime.bos] is prepended + * when the prompt does not already start with it, prompt tokens are ingested one per forward, + * and [onToken] fires only for generated tokens. + */ +@OptIn(ExperimentalMemoryApi::class) +public fun OptimizedLLMRuntime.generateTwoStage( + prompt: IntArray, + steps: Int, + temperature: Float, + head: BitNetPlanesTensorData, + native: BitNetStage1Kernel? = null, + candidates: Int = BitNetTwoStageDecode.DEFAULT_CANDIDATES, + random: Random = Random.Default, + onToken: (Int) -> Unit, +) { + require(steps > 0) { "steps must be > 0" } + val fullPrompt = when { + prompt.isEmpty() -> intArrayOf(bos) + prompt[0] != bos -> intArrayOf(bos) + prompt + else -> prompt + } + + val k = candidates.coerceIn(1, head.rows) + var token = fullPrompt[0] + var pos = 0 + var generated = 0 + while (generated < steps) { + val next = if (pos + 1 < fullPrompt.size) { + forwardHidden(token) + fullPrompt[pos + 1] + } else { + // The hidden tensor lives in the per-step forward scope — copy out before the + // next forwardHidden recycles it, exactly like logits in the stock loop. + val hidden = forwardHidden(token).data.copyToFloatArray() + val rescored = BitNetTwoStageDecode.topK( + weight = head, hidden = hidden, k = k, maxCandidates = k, native = native, + ) + sampleFromCandidates(rescored, temperature, random) + } + if (pos + 1 >= fullPrompt.size) { + onToken(next) + generated++ + } + token = next + pos++ + } +} diff --git a/llm-inference/bitnet/src/jvmTest/kotlin/sk/ainet/models/bitnet/BitNetPackedGgufLoadTest.kt b/llm-inference/bitnet/src/jvmTest/kotlin/sk/ainet/models/bitnet/BitNetPackedGgufLoadTest.kt index facf7bb5..e9673f38 100644 --- a/llm-inference/bitnet/src/jvmTest/kotlin/sk/ainet/models/bitnet/BitNetPackedGgufLoadTest.kt +++ b/llm-inference/bitnet/src/jvmTest/kotlin/sk/ainet/models/bitnet/BitNetPackedGgufLoadTest.kt @@ -329,4 +329,51 @@ class BitNetPackedGgufLoadTest { file.delete() } } + + // ---- two-stage generation loop (transformers#358) ------------------------------------ + + @Test + fun twoStageGreedyDecodeMatchesTheFullMatmulLoop() { + val file = buildFile(tied = true) + try { + val prompt = intArrayOf(1, 7, 3) + val steps = 8 + fun runtimeFor(model: Module) = + OptimizedLLMRuntime(model, ctx, OptimizedLLMMode.DIRECT, FP32::class) + + // Full-vocab matmul loop (the stock generate path). + val fullTokens = mutableListOf() + runtimeFor(loadPacked(file)).generate(prompt, steps, temperature = 0f) { fullTokens += it } + + // Two-stage, portable stage-1: trunk-only forwards + candidate rescoring. + val model2 = loadPacked(file) + val head2 = bitnetPlanesHead(model2) + assertTrue(head2 != null, "planes head must be discoverable on the loaded model") + val twoStageTokens = mutableListOf() + runtimeFor(model2).generateTwoStage(prompt, steps, temperature = 0f, head = head2) { + twoStageTokens += it + } + assertTrue( + fullTokens == twoStageTokens, + "greedy two-stage decode must equal the full-matmul loop: $fullTokens vs $twoStageTokens", + ) + + // Two-stage with the fused native stage-1, when the bundled library resolves. + if (NativeTernaryLmheadKernel.isAvailable()) { + val model3 = loadPacked(file) + val head3 = bitnetPlanesHead(model3)!! + val nativeTokens = mutableListOf() + runtimeFor(model3).generateTwoStage( + prompt, steps, temperature = 0f, head = head3, + native = BitNetStage1Kernel(NativeTernaryLmheadKernel::lmheadStage1), + ) { nativeTokens += it } + assertTrue( + fullTokens == nativeTokens, + "greedy two-stage (native stage-1) must equal the full loop: $fullTokens vs $nativeTokens", + ) + } + } finally { + file.delete() + } + } } diff --git a/llm-inference/bitnet/src/jvmTest/kotlin/sk/ainet/models/bitnet/BitNetTwoStageDecodeTest.kt b/llm-inference/bitnet/src/jvmTest/kotlin/sk/ainet/models/bitnet/BitNetTwoStageDecodeTest.kt index 7f40c812..d794a6e5 100644 --- a/llm-inference/bitnet/src/jvmTest/kotlin/sk/ainet/models/bitnet/BitNetTwoStageDecodeTest.kt +++ b/llm-inference/bitnet/src/jvmTest/kotlin/sk/ainet/models/bitnet/BitNetTwoStageDecodeTest.kt @@ -62,6 +62,30 @@ class BitNetTwoStageDecodeTest { } } + @Test + fun nativeStage1MatchesTheKotlinReference() { + // The fused lmhead_stage1 kernel and the portable loop implement the same contract + // (#358); skip quietly where the bundled native library is not available. + if (!sk.ainet.exec.kernel.NativeTernaryLmheadKernel.isAvailable()) return + val n = 48; val k = 32 + val w = weight(n, k, seed = 13) + val rng = Random(17) + repeat(3) { trial -> + val h = FloatArray(k) { rng.nextFloat() - 0.5f } + val ref = BitNetTwoStageDecode.stage1Scores(w, h) + val native = BitNetTwoStageDecode.stage1Scores( + w, h, + native = BitNetStage1Kernel(sk.ainet.exec.kernel.NativeTernaryLmheadKernel::lmheadStage1), + ) + for (r in 0 until n) { + assertTrue( + abs(ref[r] - native[r]) <= 1e-4f * maxOf(1f, abs(ref[r])), + "trial $trial row $r: native ${native[r]} vs reference ${ref[r]}", + ) + } + } + } + @Test fun sampleFromCandidatesIsGreedyAtZeroTemperatureAndStaysInTheList() { val candidates = listOf(ScoredToken(3, 1.0f), ScoredToken(17, 2.5f), ScoredToken(9, -0.5f))