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 @@ -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
Expand Down Expand Up @@ -230,6 +234,10 @@ fun main(args: Array<String>) {
// 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<FP32> = if (modelInfo.family == ModelFamily.GEMMA) {
println("Loading Gemma GGUF model from $modelPath via gemmaNetwork() + OptimizedLLMRuntime (engine loader, keep-packed, mapped)...")
if (cliArgs.contextLength != null) {
Expand Down Expand Up @@ -266,6 +274,13 @@ fun main(args: Array<String>) {
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,
Expand Down Expand Up @@ -363,9 +378,30 @@ fun main(args: Array<String>) {
println("---")
print(promptText)

val twoStageHead = bitnetTwoStageHead
@Suppress("UNCHECKED_CAST")
val twoStageRuntime = if (twoStageHead != null) runtime as? OptimizedLLMRuntime<FP32> 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

Expand Down
6 changes: 6 additions & 0 deletions llm-core/api/jvm/llm-core.api
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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 {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -173,6 +173,45 @@ public class OptimizedLLMRuntime<T : DType>(
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<T, Float> {
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
Expand Down
17 changes: 14 additions & 3 deletions llm-inference/bitnet/api/jvm/bitnet.api
Original file line number Diff line number Diff line change
Expand Up @@ -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
}

Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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<ScoredToken> {
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
Expand Down
Original file line number Diff line number Diff line change
@@ -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<FP32, Float>): BitNetPlanesTensorData? {
val head = model.modules.lastOrNull() ?: return null
@Suppress("UNCHECKED_CAST")
val params = (head as? ModuleParameters<FP32, Float>)?.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<FP32>.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++
}
}
Loading
Loading