diff --git a/.github/workflows/smoke-reference.yml b/.github/workflows/smoke-reference.yml index a2c7480b..3dac4285 100644 --- a/.github/workflows/smoke-reference.yml +++ b/.github/workflows/smoke-reference.yml @@ -39,6 +39,10 @@ on: description: "Direct URL to Apertus-8B-Instruct-2509-Q4_K_S.gguf (~4.6 GB). Enables the Apertus golden-token parity gate (QK-norm + xIELU + ungated FFN). Leave blank to skip." required: false default: "" + gemma3n_gguf_url: + description: "Direct URL to gemma-3n-E2B-it-Q4_K_M.gguf (~3.0 GB). Enables the Gemma 3n golden-token parity gate on the DSL lane (AltUp + Laurel + sparsity + PLE + shared KV; needs a large-memory runner: 20g test heap). Leave blank to skip." + required: false + default: "" gemma4_safetensors_dir_url: description: "Direct URL to a tar.gz containing the Gemma-4 E2B SafeTensors checkpoint directory. Leave blank to skip the kgemma test." required: false @@ -110,6 +114,18 @@ jobs: echo "APERTUS_GGUF_PATH=$RUNNER_TEMP/models/apertus/Apertus-8B-Instruct-2509-Q4_K_S.gguf" >> "$GITHUB_ENV" # The 8B parity gate needs more than the module's 6g default test heap. echo "APERTUS_HEAP_ARG=-PapertusTestMaxHeap=12g" >> "$GITHUB_ENV" + + - name: Stage Gemma 3n E2B GGUF + if: inputs.gemma3n_gguf_url != '' + env: + URL: ${{ inputs.gemma3n_gguf_url }} + run: | + set -euo pipefail + mkdir -p "$RUNNER_TEMP/models/gemma3n" + curl -fsSL "$URL" -o "$RUNNER_TEMP/models/gemma3n/gemma-3n-E2B-it-Q4_K_M.gguf" + echo "GEMMA3N_E2B_GGUF=$RUNNER_TEMP/models/gemma3n/gemma-3n-E2B-it-Q4_K_M.gguf" >> "$GITHUB_ENV" + # The E2B parity gate self-skips below 16 GB test heap. + echo "GEMMA3N_HEAP_ARG=-PgemmaTestMaxHeap=20g" >> "$GITHUB_ENV" if: inputs.gemma4_safetensors_dir_url != '' env: URL: ${{ inputs.gemma4_safetensors_dir_url }} @@ -149,6 +165,7 @@ jobs: -Dorg.gradle.configuration-cache=true \ -PsmokeReference -PincludeIntegration \ ${APERTUS_HEAP_ARG:-} \ + ${GEMMA3N_HEAP_ARG:-} \ test - name: Disk space (after run) diff --git a/CHANGELOG.md b/CHANGELOG.md index f8fee006..e9af4674 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,33 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added — Gemma 3n runs on the DSL path, parity-gated (#377) + +- **`gemma3nNetwork()` + `Gemma3nModel`** — the full Gemma 3n text architecture declared + in the DSL, faithful to HF `modeling_gemma3n.py`: **AltUp** (four parallel hidden + streams with the tanh modality router; `Gemma3nAltUpBlock` per layer, + `Gemma3nAltUpGlobals` for the magnitude-renormed stream init/merge), **Laurel**, + **Gaussian-top-k activation sparsity** on the first ten layers (driven by the GGUF's + precomputed per-layer std multipliers; `-inf` = off), **PLE feeding the non-active + streams** (reusing the gemma-4 lane's `PerLayerEmbedding` — the math is identical), + per-type **shared KV** for the last ten layers, hybrid sliding/global attention with + dual RoPE bases, q/k-norm + parameterless v-norm, attention scale 1.0. All math goes + through `ctx.ops`, so the model is traceable for the StableHLO → IREE mobile path. +- **The hand-rolled `Gemma3nRuntime` was never faithful to real checkpoints**: it loaded + the PLE tensors but never applied them, had no Laurel, ignored the AltUp router, and + its `E2B_DEFAULT` config claimed AltUp/sparsity were E4B-only — the real E2B GGUF has + `altup.num_inputs=4` and first-10-layer sparsity. The GGUF CLI paths (kgemma, unified + skainet-cli) now route gemma3n through the DSL lane; SafeTensors stays on the legacy + runtime until the DSL grows that leg. +- **`Gemma3nGoldenTokenParityTest`** (#346 gate, the last ungated generative family): + full 32-step greedy text equality vs mainline llama.cpp b10621 on + `gemma-3n-E2B-it-Q4_K_M.gguf`, on the exact CLI path — engine loading stays + packed/MAPPED (the PLE table row-dequants on demand). Wired into the smoke-reference + tier (`gemma3n_gguf_url` + 20g heap arg); `smoke-models.json` gains a Gemma3n-E2B row. + Metadata parsing now reads the real llama.cpp GGUF keys (`sliding_window_pattern` + booleans, per-layer `activation_sparsity_scale`, `rope.freq_base` fallback, + `rms_norm_eps`, per-layer `feed_forward_length`). + ### Fixed — Qwen tool calling follows the official Qwen3 chat template - **`QwenChatTemplate` rewritten against the official Qwen3 `chat_template`** (verified diff --git a/llm-apps/skainet-cli/build.gradle.kts b/llm-apps/skainet-cli/build.gradle.kts index 814b68af..d0638747 100644 --- a/llm-apps/skainet-cli/build.gradle.kts +++ b/llm-apps/skainet-cli/build.gradle.kts @@ -25,6 +25,7 @@ dependencies { implementation(project(":llm-inference:qwen")) implementation(project(":llm-inference:bitnet")) implementation(project(":llm-inference:gemma")) + implementation(project(":llm-inference:gemma3n")) implementation(project(":llm-inference:apertus")) // SKaiNET core libraries 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 f080686d..0fabe729 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 @@ -242,26 +242,34 @@ fun main(args: Array) { val runtime: InferenceRuntime = if (modelInfo.family == ModelFamily.GEMMA) { // ModelFamily.GEMMA claims every gemma* architecture, but this DSL lane serves - // gemma3/gemma4 only (#376): 3n needs the hand-rolled runtime (AltUp/PLE/activation - // sparsity — kgemma CLI, split tracked in #377), and gemma2 has no supported path. - when (modelInfo.architecture) { - "gemma3n" -> error( - "Gemma 3n is not supported by the unified CLI's DSL lane — use the kgemma " + - "CLI (:llm-runtime:kgemma), which carries its hand-rolled runtime (#377).", - ) - "gemma2", "gemma" -> error( - "Architecture '${modelInfo.architecture}' has no supported path — the Gemma " + - "lane serves gemma3/gemma4 checkpoints (#376).", - ) - } - println("Loading Gemma GGUF model from $modelPath via gemmaNetwork() + OptimizedLLMRuntime (engine loader, keep-packed, mapped)...") - if (cliArgs.contextLength != null) { - println(" --context flag currently ignored on the Gemma path; uses model default capped to 4096.") + // gemma3/gemma4/gemma3n (#376, #377): gemma3n runs its own DSL lane + // (gemma3nNetwork() — AltUp/Laurel/sparsity/PLE, parity-gated vs llama.cpp); + // gemma2 has no supported path. + if (modelInfo.architecture == "gemma3n") { + println("Loading Gemma 3n GGUF model from $modelPath via gemma3nNetwork() + OptimizedLLMRuntime (engine loader, keep-packed, mapped)...") + val model3n = kotlinx.coroutines.runBlocking { + sk.ainet.models.gemma3n.Gemma3nNetworkLoader.fromGguf( + ctx, + { JvmRandomAccessSource.open(modelPath.toString()) }, + ) + } + OptimizedLLMRuntime(model3n, ctx, OptimizedLLMMode.DIRECT, FP32::class) + } else { + when (modelInfo.architecture) { + "gemma2", "gemma" -> error( + "Architecture '${modelInfo.architecture}' has no supported path — the Gemma " + + "lane serves gemma3/gemma4/gemma3n checkpoints (#376, #377).", + ) + } + println("Loading Gemma GGUF model from $modelPath via gemmaNetwork() + OptimizedLLMRuntime (engine loader, keep-packed, mapped)...") + if (cliArgs.contextLength != null) { + println(" --context flag currently ignored on the Gemma path; uses model default capped to 4096.") + } + val model = GemmaNetworkLoader.fromGguf( + randomAccessProvider = { JvmRandomAccessSource.open(modelPath.toString()) } + ).load(ctx) + OptimizedLLMRuntime(model, ctx, OptimizedLLMMode.DIRECT, FP32::class) } - val model = GemmaNetworkLoader.fromGguf( - randomAccessProvider = { JvmRandomAccessSource.open(modelPath.toString()) } - ).load(ctx) - OptimizedLLMRuntime(model, ctx, OptimizedLLMMode.DIRECT, FP32::class) } else if (modelInfo.family == ModelFamily.APERTUS) { println("Loading Apertus GGUF model from $modelPath via apertusNetwork() + OptimizedLLMRuntime (engine loader, keep-packed, mapped)...") if (cliArgs.contextLength != null) { diff --git a/llm-inference/gemma3n/api/jvm/gemma3n.api b/llm-inference/gemma3n/api/jvm/gemma3n.api index 610462e0..1d0c8d21 100644 --- a/llm-inference/gemma3n/api/jvm/gemma3n.api +++ b/llm-inference/gemma3n/api/jvm/gemma3n.api @@ -49,6 +49,28 @@ public abstract interface class sk/ainet/models/gemma3n/AttentionBackend { public abstract fun reset ()V } +public final class sk/ainet/models/gemma3n/Gemma3nAltUpBlock : sk/ainet/lang/nn/Module, sk/ainet/lang/nn/topology/ModuleParameters { + public fun (IIIFLkotlin/reflect/KClass;Ljava/lang/String;)V + public synthetic fun (IIIFLkotlin/reflect/KClass;Ljava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public final fun correct (Ljava/util/List;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/context/ExecutionContext;)Ljava/util/List; + public fun getModules ()Ljava/util/List; + public fun getName ()Ljava/lang/String; + public fun getParams ()Ljava/util/List; + public final fun getRouterNorm ()Lsk/ainet/lang/nn/normalization/RMSNormalization; + public final fun predict (Ljava/util/List;Lsk/ainet/context/ExecutionContext;)Ljava/util/List; + public final fun scaleCorrectedOutput (Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/context/ExecutionContext;)Lsk/ainet/lang/tensor/Tensor; +} + +public final class sk/ainet/models/gemma3n/Gemma3nAltUpGlobals : sk/ainet/lang/nn/Module, sk/ainet/lang/nn/topology/ModuleParameters { + public fun (IILkotlin/reflect/KClass;Ljava/lang/String;)V + public synthetic fun (IILkotlin/reflect/KClass;Ljava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public fun getModules ()Ljava/util/List; + public fun getName ()Ljava/lang/String; + public fun getParams ()Ljava/util/List; + public final fun initStreams (Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/context/ExecutionContext;)Ljava/util/List; + public final fun mergeStreams (Ljava/util/List;Lsk/ainet/context/ExecutionContext;)Lsk/ainet/lang/tensor/Tensor; +} + public final class sk/ainet/models/gemma3n/Gemma3nAttentionBackend : sk/ainet/models/gemma3n/AttentionBackend { public fun (Lsk/ainet/context/ExecutionContext;Lsk/ainet/models/gemma3n/Gemma3nRuntimeWeights;Lkotlin/reflect/KClass;Lsk/ainet/models/gemma3n/Gemma3nConfig;Lsk/ainet/models/gemma3n/Gemma3nKvCache;)V public synthetic fun (Lsk/ainet/context/ExecutionContext;Lsk/ainet/models/gemma3n/Gemma3nRuntimeWeights;Lkotlin/reflect/KClass;Lsk/ainet/models/gemma3n/Gemma3nConfig;Lsk/ainet/models/gemma3n/Gemma3nKvCache;ILkotlin/jvm/internal/DefaultConstructorMarker;)V @@ -123,6 +145,11 @@ public final class sk/ainet/models/gemma3n/Gemma3nConfigParser { public final fun parseFromJson (Ljava/lang/String;)Lsk/ainet/models/gemma3n/Gemma3nModelMetadata; } +public final class sk/ainet/models/gemma3n/Gemma3nGGUFNameResolver : sk/ainet/io/weights/WeightNameResolver { + public fun ()V + public fun resolve (Ljava/lang/String;Ljava/lang/String;)Ljava/lang/String; +} + public abstract interface class sk/ainet/models/gemma3n/Gemma3nKvCache : sk/ainet/apps/llm/KvCache { } @@ -130,6 +157,16 @@ public final class sk/ainet/models/gemma3n/Gemma3nKvCacheKt { public static final fun createOptimalGemma3nKvCache (Lsk/ainet/models/gemma3n/Gemma3nConfig;I)Lsk/ainet/models/gemma3n/Gemma3nKvCache; } +public final class sk/ainet/models/gemma3n/Gemma3nLaurelBlock : sk/ainet/lang/nn/Module { + public fun (IIFLkotlin/reflect/KClass;Ljava/lang/String;)V + public synthetic fun (IIFLkotlin/reflect/KClass;Ljava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public final fun getLinearLeft ()Lsk/ainet/lang/nn/transformer/VoidDense; + public final fun getLinearRight ()Lsk/ainet/lang/nn/transformer/VoidDense; + public fun getModules ()Ljava/util/List; + public fun getName ()Ljava/lang/String; + public final fun getPostNorm ()Lsk/ainet/lang/nn/normalization/RMSNormalization; +} + public final class sk/ainet/models/gemma3n/Gemma3nLayerWeights { public fun (Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/models/gemma3n/AltUpLayerWeights;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;)V public synthetic fun (Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/models/gemma3n/AltUpLayerWeights;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;ILkotlin/jvm/internal/DefaultConstructorMarker;)V @@ -184,14 +221,43 @@ public final class sk/ainet/models/gemma3n/Gemma3nLayerWeights { public fun toString ()Ljava/lang/String; } +public final class sk/ainet/models/gemma3n/Gemma3nModel : sk/ainet/lang/nn/Module { + public fun (Lsk/ainet/lang/nn/layers/EmbeddingAdapter;Lsk/ainet/models/gemma/PerLayerEmbedding;Lsk/ainet/models/gemma3n/Gemma3nAltUpGlobals;Ljava/util/List;Lsk/ainet/lang/nn/normalization/RMSNormalization;Lsk/ainet/lang/nn/transformer/VoidDense;Lkotlin/reflect/KClass;IFLjava/lang/String;)V + public synthetic fun (Lsk/ainet/lang/nn/layers/EmbeddingAdapter;Lsk/ainet/models/gemma/PerLayerEmbedding;Lsk/ainet/models/gemma3n/Gemma3nAltUpGlobals;Ljava/util/List;Lsk/ainet/lang/nn/normalization/RMSNormalization;Lsk/ainet/lang/nn/transformer/VoidDense;Lkotlin/reflect/KClass;IFLjava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public final fun getActiveIdx ()I + public final fun getAltupGlobals ()Lsk/ainet/models/gemma3n/Gemma3nAltUpGlobals; + public final fun getBlocks ()Ljava/util/List; + public final fun getDtype ()Lkotlin/reflect/KClass; + public final fun getEmbedScale ()F + public final fun getLmHead ()Lsk/ainet/lang/nn/transformer/VoidDense; + public fun getModules ()Ljava/util/List; + public fun getName ()Ljava/lang/String; + public final fun getOutputNorm ()Lsk/ainet/lang/nn/normalization/RMSNormalization; + public final fun getPle ()Lsk/ainet/models/gemma/PerLayerEmbedding; + public final fun getTokenEmbedding ()Lsk/ainet/lang/nn/layers/EmbeddingAdapter; +} + +public final class sk/ainet/models/gemma3n/Gemma3nModel$BlockRefs { + public fun (Lsk/ainet/lang/nn/normalization/RMSNormalization;Lsk/ainet/lang/nn/transformer/MultiHeadAttention;Lsk/ainet/lang/nn/normalization/RMSNormalization;Lsk/ainet/lang/nn/normalization/RMSNormalization;Lsk/ainet/models/gemma3n/Gemma3nSparseGeGluFFN;Lsk/ainet/lang/nn/normalization/RMSNormalization;Lsk/ainet/models/gemma3n/Gemma3nLaurelBlock;Lsk/ainet/models/gemma3n/Gemma3nAltUpBlock;Lsk/ainet/models/gemma3n/Gemma3nPerLayerApply;)V + public final fun getAltup ()Lsk/ainet/models/gemma3n/Gemma3nAltUpBlock; + public final fun getAttnNorm ()Lsk/ainet/lang/nn/normalization/RMSNormalization; + public final fun getFfn ()Lsk/ainet/models/gemma3n/Gemma3nSparseGeGluFFN; + public final fun getFfnNorm ()Lsk/ainet/lang/nn/normalization/RMSNormalization; + public final fun getLaurel ()Lsk/ainet/models/gemma3n/Gemma3nLaurelBlock; + public final fun getMha ()Lsk/ainet/lang/nn/transformer/MultiHeadAttention; + public final fun getPerLayer ()Lsk/ainet/models/gemma3n/Gemma3nPerLayerApply; + public final fun getPostAttnNorm ()Lsk/ainet/lang/nn/normalization/RMSNormalization; + public final fun getPostFfwNorm ()Lsk/ainet/lang/nn/normalization/RMSNormalization; +} + public final class sk/ainet/models/gemma3n/Gemma3nModelMetadata { public static final field Companion Lsk/ainet/models/gemma3n/Gemma3nModelMetadata$Companion; public static final field DEFAULT_KV_SHARED_LAYERS I public static final field DEFAULT_ROPE_BASE_GLOBAL F public static final field DEFAULT_ROPE_BASE_LOCAL F public static final field DEFAULT_SLIDING_WINDOW I - public fun (Ljava/lang/String;IIIIIILjava/util/List;IIIFFILjava/util/List;IILjava/util/List;F)V - public synthetic fun (Ljava/lang/String;IIIIIILjava/util/List;IIIFFILjava/util/List;IILjava/util/List;FILkotlin/jvm/internal/DefaultConstructorMarker;)V + public fun (Ljava/lang/String;IIIIIILjava/util/List;IIIFFILjava/util/List;IILjava/util/List;FFLjava/util/List;)V + public synthetic fun (Ljava/lang/String;IIIIIILjava/util/List;IIIFFILjava/util/List;IILjava/util/List;FFLjava/util/List;ILkotlin/jvm/internal/DefaultConstructorMarker;)V public final fun component1 ()Ljava/lang/String; public final fun component10 ()I public final fun component11 ()I @@ -204,6 +270,8 @@ public final class sk/ainet/models/gemma3n/Gemma3nModelMetadata { public final fun component18 ()Ljava/util/List; public final fun component19 ()F public final fun component2 ()I + public final fun component20 ()F + public final fun component21 ()Ljava/util/List; public final fun component3 ()I public final fun component4 ()I public final fun component5 ()I @@ -211,11 +279,12 @@ public final class sk/ainet/models/gemma3n/Gemma3nModelMetadata { public final fun component7 ()I public final fun component8 ()Ljava/util/List; public final fun component9 ()I - public final fun copy (Ljava/lang/String;IIIIIILjava/util/List;IIIFFILjava/util/List;IILjava/util/List;F)Lsk/ainet/models/gemma3n/Gemma3nModelMetadata; - public static synthetic fun copy$default (Lsk/ainet/models/gemma3n/Gemma3nModelMetadata;Ljava/lang/String;IIIIIILjava/util/List;IIIFFILjava/util/List;IILjava/util/List;FILjava/lang/Object;)Lsk/ainet/models/gemma3n/Gemma3nModelMetadata; + public final fun copy (Ljava/lang/String;IIIIIILjava/util/List;IIIFFILjava/util/List;IILjava/util/List;FFLjava/util/List;)Lsk/ainet/models/gemma3n/Gemma3nModelMetadata; + public static synthetic fun copy$default (Lsk/ainet/models/gemma3n/Gemma3nModelMetadata;Ljava/lang/String;IIIIIILjava/util/List;IIIFFILjava/util/List;IILjava/util/List;FFLjava/util/List;ILjava/lang/Object;)Lsk/ainet/models/gemma3n/Gemma3nModelMetadata; public fun equals (Ljava/lang/Object;)Z public final fun getActivationSparsityPattern ()Ljava/util/List; public final fun getActivationSparsityScale ()F + public final fun getActivationSparsityScales ()Ljava/util/List; public final fun getAltupActiveIdx ()I public final fun getArchitecture ()Ljava/lang/String; public final fun getBlockCount ()I @@ -233,6 +302,7 @@ public final class sk/ainet/models/gemma3n/Gemma3nModelMetadata { public final fun getLayerType (I)Lsk/ainet/models/gemma/LayerType; public final fun getNumAltupInputs ()I public final fun getPerLayerEmbeddingLength ()I + public final fun getRmsNormEps ()F public final fun getRopeBase (I)F public final fun getRopeBaseGlobal ()F public final fun getRopeBaseLocal ()F @@ -240,6 +310,7 @@ public final class sk/ainet/models/gemma3n/Gemma3nModelMetadata { public final fun getVocabSize ()I public fun hashCode ()I public final fun isKvShared (I)Z + public final fun sparsityScaleFor (I)Ljava/lang/Float; public fun toString ()Ljava/lang/String; } @@ -247,6 +318,29 @@ public final class sk/ainet/models/gemma3n/Gemma3nModelMetadata$Companion { public final fun getDEFAULT_LAYER_PATTERN ()Ljava/util/List; } +public final class sk/ainet/models/gemma3n/Gemma3nNetworkDefKt { + public static final field LAUREL_RANK I + public static final fun gemma3nNetwork (Lsk/ainet/models/gemma3n/Gemma3nModelMetadata;Lkotlin/reflect/KClass;II)Lsk/ainet/lang/nn/Module; + public static synthetic fun gemma3nNetwork$default (Lsk/ainet/models/gemma3n/Gemma3nModelMetadata;Lkotlin/reflect/KClass;IIILjava/lang/Object;)Lsk/ainet/lang/nn/Module; +} + +public final class sk/ainet/models/gemma3n/Gemma3nNetworkLoader { + public static final field INSTANCE Lsk/ainet/models/gemma3n/Gemma3nNetworkLoader; + public final fun fromWeights (Lsk/ainet/context/ExecutionContext;Lsk/ainet/models/gemma3n/Gemma3nWeights;Lkotlin/reflect/KClass;Ljava/lang/Integer;Z)Lsk/ainet/lang/nn/Module; + public static synthetic fun fromWeights$default (Lsk/ainet/models/gemma3n/Gemma3nNetworkLoader;Lsk/ainet/context/ExecutionContext;Lsk/ainet/models/gemma3n/Gemma3nWeights;Lkotlin/reflect/KClass;Ljava/lang/Integer;ZILjava/lang/Object;)Lsk/ainet/lang/nn/Module; +} + +public final class sk/ainet/models/gemma3n/Gemma3nPerLayerApply : sk/ainet/lang/nn/Module { + public fun (IIFLkotlin/reflect/KClass;Ljava/lang/String;)V + public synthetic fun (IIFLkotlin/reflect/KClass;Ljava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public final fun computeDelta (Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/context/ExecutionContext;)Lsk/ainet/lang/tensor/Tensor; + public final fun getInpGate ()Lsk/ainet/lang/nn/transformer/VoidDense; + public fun getModules ()Ljava/util/List; + public fun getName ()Ljava/lang/String; + public final fun getPostNorm ()Lsk/ainet/lang/nn/normalization/RMSNormalization; + public final fun getProj ()Lsk/ainet/lang/nn/transformer/VoidDense; +} + public final class sk/ainet/models/gemma3n/Gemma3nRuntime : sk/ainet/apps/llm/DecoderRuntime { public static final field BOS_TOKEN I public fun (Lsk/ainet/context/ExecutionContext;Lsk/ainet/models/gemma3n/Gemma3nRuntimeWeights;Lsk/ainet/models/gemma3n/AttentionBackend;Lkotlin/reflect/KClass;Lsk/ainet/models/gemma3n/Gemma3nConfig;FLkotlin/random/Random;)V @@ -314,6 +408,17 @@ public final class sk/ainet/models/gemma3n/Gemma3nSafeTensorsWeightLoaderKt { public static final fun loadGemma3nRuntimeWeightsFromSafeTensors (Lsk/ainet/context/ExecutionContext;Ljava/lang/String;Lkotlin/reflect/KClass;Lkotlin/coroutines/Continuation;)Ljava/lang/Object; } +public final class sk/ainet/models/gemma3n/Gemma3nSparseGeGluFFN : sk/ainet/lang/nn/Module { + public fun (IIFLkotlin/reflect/KClass;Ljava/lang/String;)V + public synthetic fun (IIFLkotlin/reflect/KClass;Ljava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public final fun getDown ()Lsk/ainet/lang/nn/transformer/VoidDense; + public final fun getGate ()Lsk/ainet/lang/nn/transformer/VoidDense; + public fun getModules ()Ljava/util/List; + public fun getName ()Ljava/lang/String; + public final fun getSparsityEnabled ()Z + public final fun getUp ()Lsk/ainet/lang/nn/transformer/VoidDense; +} + public final class sk/ainet/models/gemma3n/Gemma3nTensorNames { public static final field ALTUP_PROJ Ljava/lang/String; public static final field ALTUP_UNEMBD_PROJ Ljava/lang/String; diff --git a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nDslModules.kt b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nDslModules.kt new file mode 100644 index 00000000..7eabeb5b --- /dev/null +++ b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nDslModules.kt @@ -0,0 +1,313 @@ +package sk.ainet.models.gemma3n + +import sk.ainet.context.ExecutionContext +import sk.ainet.lang.nn.Module +import sk.ainet.lang.nn.normalization.RMSNormalization +import sk.ainet.lang.nn.topology.ModuleParameter +import sk.ainet.lang.nn.topology.ModuleParameters +import sk.ainet.lang.nn.transformer.VoidDense +import sk.ainet.lang.nn.transformer.linearProject +import sk.ainet.lang.tensor.Shape +import sk.ainet.lang.tensor.Tensor +import sk.ainet.lang.tensor.VoidOpsTensor +import sk.ainet.lang.tensor.data.TensorData +import sk.ainet.lang.types.DType +import kotlin.reflect.KClass + +/* + * DSL modules for the Gemma 3n-specific machinery (the #377 DSL migration): AltUp, Laurel, + * activation-sparsity FFN and the per-layer-input application. All math goes through + * `ctx.ops` so the modules are traceable for the StableHLO → IREE export path. The + * reference implementation is HF `transformers` `modeling_gemma3n.py` (verified against + * the installed 5.x source); working tensors are rank-2 `[seq, hidden]`. + */ + +@Suppress("UNCHECKED_CAST") +internal fun voidParam(name: String, shape: Shape, dtype: KClass?): ModuleParameter = + ModuleParameter.WeightParameter( + name, + VoidOpsTensor( + object : TensorData { + override val shape: Shape = shape + override fun get(vararg indices: Int): V = 0.0f as V + override fun set(vararg indices: Int, value: V) {} + }, + (dtype ?: Any::class) as KClass, + ), + ) + +/** + * Per-layer AltUp (Alternating Updates) block — HF `Gemma3nTextAltUp`. + * + * Maintains `numInputs` parallel hidden streams; only the active one runs the expensive + * transformer sub-layers, the rest are predicted/corrected via a learned router: + * + * ``` + * modalities(x) = tanh( modality_router( router_norm(x) * 1/hidden ) ) + * predict: coefs = prediction_coefs(modalities) # [S, n²] + * pred_i = h_i + Σ_j coefs[:, i·n+j] ⊙ h_j + * correct: coefs = correction_coefs(modalities(activated)) + 1 + * innovation = activated − pred_active + * corr_i = pred_i + coefs[:, i] ⊙ innovation + * scale_corrected_output(x) = x * correct_output_scale # [hidden] + * ``` + */ +public class Gemma3nAltUpBlock( + private val hiddenSize: Int, + private val numInputs: Int, + private val activeIdx: Int, + rmsEps: Float, + private val dtype: KClass? = null, + override val name: String = "altup", +) : Module(), ModuleParameters { + + /** `router_norm` — scale-full RMSNorm over hidden, plain (non-unit-offset) weight. */ + public val routerNorm: RMSNormalization = RMSNormalization( + intArrayOf(hiddenSize), rmsEps.toDouble(), unitOffset = false, name = "$name.altup_router_norm", dtype = dtype, + ) + + override val params: List> = listOf( + voidParam("$name.altup_router.weight", Shape(numInputs, hiddenSize), dtype), + voidParam("$name.altup_predict_coef.weight", Shape(numInputs * numInputs, numInputs), dtype), + voidParam("$name.altup_correct_coef.weight", Shape(numInputs, numInputs), dtype), + voidParam("$name.altup_correct_scale.weight", Shape(hiddenSize), dtype), + ) + + override val modules: List> = listOf(routerNorm) + + override fun onForward(input: Tensor, ctx: ExecutionContext): Tensor = input + + private fun modalities(x: Tensor, ctx: ExecutionContext): Tensor { + val ops = ctx.ops + val normed = routerNorm.forward(x, ctx) + val scaled = ops.mulScalar(normed, 1.0f / hiddenSize) + return ops.tanh(linearProject(ops, scaled, params[0].value)) // [S, n] + } + + /** One learned scalar column `[S, 1]` broadcast-multiplied over `[S, H]`. */ + private fun scaleBy(coefs: Tensor, col: Int, x: Tensor, ctx: ExecutionContext): Tensor = + ctx.ops.multiply(x, ctx.ops.narrow(coefs, dim = 1, start = col, length = 1)) + + public fun predict(streams: List>, ctx: ExecutionContext): List> { + val ops = ctx.ops + val m = modalities(streams[activeIdx], ctx) + val coefs = linearProject(ops, m, params[1].value) // [S, n²] + return List(numInputs) { i -> + var pred = streams[i] + for (j in 0 until numInputs) { + pred = ops.add(pred, scaleBy(coefs, i * numInputs + j, streams[j], ctx)) + } + pred + } + } + + public fun correct( + predictions: List>, + activated: Tensor, + ctx: ExecutionContext, + ): List> { + val ops = ctx.ops + val m = modalities(activated, ctx) + val coefs = ops.addScalar(linearProject(ops, m, params[2].value), 1.0f) // [S, n] + val innovation = ops.subtract(activated, predictions[activeIdx]) + return List(numInputs) { i -> + ops.add(predictions[i], scaleBy(coefs, i, innovation, ctx)) + } + } + + /** `x * correct_output_scale` (element-wise over hidden). */ + public fun scaleCorrectedOutput(x: Tensor, ctx: ExecutionContext): Tensor = + ctx.ops.multiply(x, params[3].value) +} + +/** + * Model-level AltUp stream projections — HF `altup_projections` / `altup_unembed_projections`. + * The GGUF stores each set as ONE 3D tensor (`altup_proj.weight`, `altup_unembd_proj.weight`, + * logical `[numInputs-1, hidden, hidden]`); slices are narrowed out at forward time. + * + * Both directions renormalize the projected stream to the active stream's per-token RMS + * magnitude (HF: `target_magnitude / max(rms(proj), 1e-5)`). + */ +public class Gemma3nAltUpGlobals( + private val hiddenSize: Int, + private val numInputs: Int, + private val dtype: KClass? = null, + override val name: String = "altup_globals", +) : Module(), ModuleParameters { + + override val params: List> = listOf( + voidParam("$name.altup_proj.weight", Shape(numInputs - 1, hiddenSize, hiddenSize), dtype), + voidParam("$name.altup_unembd_proj.weight", Shape(numInputs - 1, hiddenSize, hiddenSize), dtype), + ) + + override val modules: List> = emptyList() + + override fun onForward(input: Tensor, ctx: ExecutionContext): Tensor = input + + /** Per-token RMS magnitude `[S, 1]`: `sqrt(mean(x², dim=-1))`. */ + private fun magnitude(x: Tensor, ctx: ExecutionContext): Tensor { + val ops = ctx.ops + val meanSq = ops.mean(ops.multiply(x, x), dim = -1) // [S] + return ops.unsqueeze(ops.sqrt(meanSq), dim = -1) // [S, 1] + } + + private fun sliceOf(param: ModuleParameter, k: Int, ctx: ExecutionContext): Tensor { + val ops = ctx.ops + // The GGUF stores the stack as ne=[hidden, hidden, numExtra] (ggml: ne2 slowest), and + // the engine surfaces the raw buffer under that ne-ordered shape — so the slice index + // is SLOWEST in memory. Reinterpret row-major as [numExtra, hidden, hidden] first, + // then narrow the leading dim; each slice's buffer is the converter's [out, in] + // row-major matrix. + val stacked = ops.reshape(param.value, Shape(numInputs - 1, hiddenSize, hiddenSize)) + val sliced = ops.narrow(stacked, dim = 0, start = k, length = 1) + return ops.reshape(sliced, Shape(hiddenSize, hiddenSize)) + } + + private fun projectRenormed( + x0mag: Tensor, + stream: Tensor, + param: ModuleParameter, + k: Int, + ctx: ExecutionContext, + ): Tensor { + val ops = ctx.ops + val proj = linearProject(ops, stream, sliceOf(param, k, ctx)) + val newMagSq = ops.mean(ops.multiply(proj, proj), dim = -1) // [S] + val newMag = ops.unsqueeze(ops.sqrt(ops.clamp(newMagSq, 1e-5f, Float.MAX_VALUE)), dim = -1) + return ops.multiply(proj, ops.divide(x0mag, newMag)) + } + + /** HF stream init: `[h0] + [renorm(altup_projections[k](h0))]`. */ + public fun initStreams(h0: Tensor, ctx: ExecutionContext): List> { + val mag = magnitude(h0, ctx) + return listOf(h0) + List(numInputs - 1) { k -> projectRenormed(mag, h0, params[0], k, ctx) } + } + + /** HF finalize: mean of `[h0] + [renorm(altup_unembed_projections[k](h_k+1))]`. */ + public fun mergeStreams(streams: List>, ctx: ExecutionContext): Tensor { + val ops = ctx.ops + val mag = magnitude(streams[0], ctx) + var acc = streams[0] + for (k in 0 until numInputs - 1) { + acc = ops.add(acc, projectRenormed(mag, streams[k + 1], params[1], k, ctx)) + } + return ops.mulScalar(acc, 1.0f / numInputs) + } +} + +/** + * Laurel (Learned Augmented Residual Layer) — HF `Gemma3nTextLaurelBlock`: + * `x + post_laurel_norm(linear_right(linear_left(x)))`. + */ +public class Gemma3nLaurelBlock( + hiddenSize: Int, + laurelRank: Int, + rmsEps: Float, + dtype: KClass? = null, + override val name: String = "laurel", +) : Module() { + + public val linearLeft: VoidDense = VoidDense("$name.laurel_l", laurelRank, hiddenSize, dtype) + public val linearRight: VoidDense = VoidDense("$name.laurel_r", hiddenSize, laurelRank, dtype) + public val postNorm: RMSNormalization = RMSNormalization( + intArrayOf(hiddenSize), rmsEps.toDouble(), unitOffset = false, name = "$name.laurel_post_norm", dtype = dtype, + ) + + override val modules: List> = listOf(linearLeft, linearRight, postNorm) + + override fun onForward(input: Tensor, ctx: ExecutionContext): Tensor { + val low = linearLeft.forward(input, ctx) + val back = linearRight.forward(low, ctx) + return ctx.ops.add(input, postNorm.forward(back, ctx)) + } +} + +/** + * Gemma 3n FFN — gelu-gated (`down(gelu(gate(x)) * up(x))`) with optional Gaussian-top-k + * activation sparsity on the gate projection (HF `Gemma3nTextMLP._gaussian_topk`): + * + * ``` + * cutoff = mean(gate, -1) + std_pop(gate, -1) * stdMultiplier + * gate = relu(gate - cutoff) + * ``` + * + * `stdMultiplier` comes precomputed per layer from the GGUF (`activation_sparsity_scale`, + * `Φ⁻¹(0.95) ≈ 1.6449` on sparse layers, `-inf` on the rest — non-finite disables the + * whole branch at build time). Std is population (unbiased=False): `sqrt(E[x²] − E[x]²)`. + */ +public class Gemma3nSparseGeGluFFN( + hiddenSize: Int, + ffnDim: Int, + private val stdMultiplier: Float, + dtype: KClass? = null, + override val name: String = "ffn", +) : Module() { + + // Param names follow the llama/HF convention the engine resolver maps to + // `blk.N.ffn_{gate,up,down}.weight`. + public val gate: VoidDense = VoidDense("$name.gate_proj", ffnDim, hiddenSize, dtype) + public val up: VoidDense = VoidDense("$name.up_proj", ffnDim, hiddenSize, dtype) + public val down: VoidDense = VoidDense("$name.down_proj", hiddenSize, ffnDim, dtype) + + public val sparsityEnabled: Boolean = stdMultiplier.isFinite() && stdMultiplier > 0f + + override val modules: List> = listOf(gate, up, down) + + override fun onForward(input: Tensor, ctx: ExecutionContext): Tensor { + val ops = ctx.ops + var g = gate.forward(input, ctx) + if (sparsityEnabled) { + val mean = ops.mean(g, dim = -1) // [S] + val meanSq = ops.mean(ops.multiply(g, g), dim = -1) // [S] + val varPop = ops.subtract(meanSq, ops.multiply(mean, mean)) + val std = ops.sqrt(ops.clamp(varPop, 0f, Float.MAX_VALUE)) + val cutoff = ops.unsqueeze( + ops.add(mean, ops.mulScalar(std, stdMultiplier)), dim = -1, + ) // [S, 1] + g = ops.relu(ops.subtract(g, cutoff)) + } + val activated = ops.gelu(g) + val upOut = up.forward(input, ctx) + return down.forward(ops.multiply(activated, upOut), ctx) + } +} + +/** + * Per-layer-input application — the tail of HF `Gemma3nTextDecoderLayer.forward`. + * Takes the (scaled) corrected active stream and this layer's `per_layer_input` slice, + * returns the DELTA that gets added to the non-active streams: + * `post_norm( proj( gelu(inp_gate(x)) ⊙ per_layer_input ) )`. + * + * The gemma-4 lane's `PerLayerInputBlockHook` applies the same transform but adds it to + * the main residual (gemma-4 has no AltUp streams); gemma3n adds it to streams `1..n-1`, + * so this module returns the delta and `Gemma3nModel` does the stream adds. + */ +public class Gemma3nPerLayerApply( + hiddenSize: Int, + perLayerDim: Int, + rmsEps: Float, + dtype: KClass? = null, + override val name: String = "per_layer_input", +) : Module() { + + public val inpGate: VoidDense = VoidDense("$name.inp_gate", perLayerDim, hiddenSize, dtype) + public val proj: VoidDense = VoidDense("$name.proj", hiddenSize, perLayerDim, dtype) + public val postNorm: RMSNormalization = RMSNormalization( + intArrayOf(hiddenSize), rmsEps.toDouble(), unitOffset = false, name = "$name.post_norm", dtype = dtype, + ) + + override val modules: List> = listOf(inpGate, proj, postNorm) + + override fun onForward(input: Tensor, ctx: ExecutionContext): Tensor = input + + public fun computeDelta( + activeCorrected: Tensor, + perLayerInput: Tensor, + ctx: ExecutionContext, + ): Tensor { + val ops = ctx.ops + val gated = ops.gelu(inpGate.forward(activeCorrected, ctx)) + val mixed = ops.multiply(gated, perLayerInput) + return postNorm.forward(proj.forward(mixed, ctx), ctx) + } +} diff --git a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nGGUFNameResolver.kt b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nGGUFNameResolver.kt new file mode 100644 index 00000000..65f1dc3f --- /dev/null +++ b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nGGUFNameResolver.kt @@ -0,0 +1,40 @@ +package sk.ainet.models.gemma3n + +import sk.ainet.io.weights.WeightNameResolver +import sk.ainet.models.gemma.GemmaGGUFNameResolver + +/** + * Resolves DSL module paths to GGUF tensor names for the Gemma 3n family: the 3n-specific + * rules (AltUp per-layer + global tensors, Laurel) matched first, everything the gemma-4 + * lane already handles (sandwich norms, PLE names, llama-standard set) delegated to + * [GemmaGGUFNameResolver]. + */ +public class Gemma3nGGUFNameResolver : WeightNameResolver { + + private val gemma = GemmaGGUFNameResolver() + + override fun resolve(modulePath: String, paramName: String): String? { + val blockPrefix = modulePath.split("/").drop(1).firstOrNull { it.startsWith("blk.") } + + // Per-layer AltUp + Laurel params are named after their GGUF tensors already — + // ".altup_router.weight" etc. — so the rule is: take the tensor-suffix + // and prefix the block. + for (suffix in BLOCK_SUFFIXES) { + if (paramName.endsWith(".$suffix.weight") || paramName == "$suffix.weight") { + return if (blockPrefix != null) "$blockPrefix.$suffix.weight" else null + } + } + // Model-level AltUp stream projections (3D tensors, no block prefix). + if (paramName.endsWith(".altup_proj.weight")) return "altup_proj.weight" + if (paramName.endsWith(".altup_unembd_proj.weight")) return "altup_unembd_proj.weight" + + return gemma.resolve(modulePath, paramName) + } + + private companion object { + val BLOCK_SUFFIXES = listOf( + "altup_router_norm", "altup_router", "altup_predict_coef", "altup_correct_coef", + "altup_correct_scale", "laurel_l", "laurel_r", "laurel_post_norm", + ) + } +} diff --git a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nModel.kt b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nModel.kt new file mode 100644 index 00000000..5a08ab0c --- /dev/null +++ b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nModel.kt @@ -0,0 +1,153 @@ +package sk.ainet.models.gemma3n + +import sk.ainet.apps.llm.HybridTransformerBlock +import sk.ainet.context.ExecutionContext +import sk.ainet.lang.nn.Module +import sk.ainet.lang.nn.layers.EmbeddingAdapter +import sk.ainet.lang.nn.normalization.RMSNormalization +import sk.ainet.lang.nn.transformer.MultiHeadAttention +import sk.ainet.lang.nn.transformer.VoidDense +import sk.ainet.lang.tensor.Tensor +import sk.ainet.lang.types.DType +import sk.ainet.models.gemma.PerLayerEmbedding +import kotlin.reflect.KClass + +/** + * Top-level Gemma 3n model — the DSL replacement for the hand-rolled `Gemma3nRuntime` + * (#377), faithful to HF `Gemma3nTextModel.forward` / `Gemma3nTextDecoderLayer.forward`. + * + * Follows the `GemmaModel` wrapper pattern (gemma-4 PLE precedent): the module tree is + * regular DSL modules (so `WeightMapper` binds every weight by name), but the forward + * orchestration is bespoke because AltUp threads `numInputs` parallel hidden streams + * through every layer — inexpressible as a plain Sequential: + * + * ``` + * h0 = embed(ids) * sqrt(hidden) + * ple = PerLayerEmbedding.compute(ids, h0) # [B, S, L, pleDim] + * streams = altupGlobals.initStreams(h0) # magnitude-renormed projections + * per layer: + * preds = altup.predict(streams) + * active = preds[activeIdx]; an = attn_norm(active) + * laurel = laurel(an) # an + norm(right(left(an))) + * attn = post_attention_norm( MHA(an) ) + * attnLaurel = ((active + attn) + laurel) / √2 + * ffw = post_ffw_norm( ffn( ffn_norm(attnLaurel) ) ) # sparsity on first layers + * streams = altup.correct(preds, attnLaurel + ffw) + * delta = perLayerApply( scale(streams[active]), ple[:, :, layer] ) + * streams[1:] += delta + * merged = altupGlobals.mergeStreams(streams) # renormed mean + * logits = lm_head( output_norm(merged) ) # tied embeddings, no softcap + * ``` + */ +public class Gemma3nModel( + public val tokenEmbedding: EmbeddingAdapter, + public val ple: PerLayerEmbedding, + public val altupGlobals: Gemma3nAltUpGlobals, + public val blocks: List>, + public val outputNorm: RMSNormalization, + public val lmHead: VoidDense, + public val dtype: KClass, + public val activeIdx: Int, + public val embedScale: Float, + override val name: String = "Gemma3nModel", +) : Module() { + + override val modules: List> = buildList { + add(tokenEmbedding) + add(ple) + add(altupGlobals) + addAll(blocks) + add(outputNorm) + add(lmHead) + } + + /** Typed handles into one block's module list (bound by construction in `gemma3nNetwork`). */ + public class BlockRefs( + public val attnNorm: RMSNormalization, + public val mha: MultiHeadAttention, + public val postAttnNorm: RMSNormalization, + public val ffnNorm: RMSNormalization, + public val ffn: Gemma3nSparseGeGluFFN, + public val postFfwNorm: RMSNormalization, + public val laurel: Gemma3nLaurelBlock, + public val altup: Gemma3nAltUpBlock, + public val perLayer: Gemma3nPerLayerApply, + ) + + @Suppress("UNCHECKED_CAST") + private fun refsFor(block: HybridTransformerBlock): BlockRefs { + val mods = block.modules + fun norm(id: String): RMSNormalization = + mods.filterIsInstance>().firstOrNull { it.name == id } + ?: error("Gemma3nModel: block ${block.name} has no RMSNorm '$id'") + return BlockRefs( + attnNorm = norm("attn_norm"), + mha = mods.filterIsInstance>().first(), + postAttnNorm = norm("post_attention_norm"), + ffnNorm = norm("ffn_norm"), + ffn = mods.filterIsInstance>().first(), + postFfwNorm = norm("post_ffw_norm"), + laurel = mods.filterIsInstance>().first(), + altup = mods.filterIsInstance>().first(), + perLayer = mods.filterIsInstance>().first(), + ) + } + + override fun onForward(input: Tensor, ctx: ExecutionContext): Tensor { + val ops = ctx.ops + val invSqrt2 = 0.70710678f + + // 1 — scaled embedding. + val rawEmbeds = tokenEmbedding.forward(input, ctx) + val h0 = if (embedScale != 1f) ops.mulScalar(rawEmbeds, embedScale) else rawEmbeds + + // 2 — per-layer inputs [B, S, L, pleDim] (identical math to gemma-4: token-identity + // gather * sqrt(pleDim), context projection * hidden^-0.5, norm, sum * 1/sqrt2). + val ids2d = if (input.rank == 1) ops.unsqueeze(input, 0) else input + val embeds3d = if (h0.rank == 2) ops.unsqueeze(h0, 0) else h0 + val perLayerInputs = ple.compute(ids2d, embeds3d, ctx, dtype) + + // 3 — AltUp stream init. + var streams = altupGlobals.initStreams(h0, ctx) + + // 4 — per-layer flow. + for ((layerIdx, block) in blocks.withIndex()) { + val r = refsFor(block) + val preds = r.altup.predict(streams, ctx) + val active = preds[activeIdx] + val an = r.attnNorm.forward(active, ctx) + val laurel = r.laurel.forward(an, ctx) + val attn = r.postAttnNorm.forward(r.mha.forward(an, ctx), ctx) + val attnGated = ops.add(active, attn) + val attnLaurel = ops.mulScalar(ops.add(attnGated, laurel), invSqrt2) + val ffw = r.postFfwNorm.forward(r.ffn.forward(r.ffnNorm.forward(attnLaurel, ctx), ctx), ctx) + val corrected = r.altup.correct(preds, ops.add(attnLaurel, ffw), ctx) + + // Per-layer input: transform the (scaled) corrected active stream and add the + // delta to the NON-active streams (HF: `corrected_predictions[1:] += ...`). + val scaledActive = r.altup.scaleCorrectedOutput(corrected[activeIdx], ctx) + val pliSlice = perLayerSlice(perLayerInputs, layerIdx, h0.rank, ctx) + val delta = r.perLayer.computeDelta(scaledActive, pliSlice, ctx) + streams = List(corrected.size) { i -> + if (i == 0) corrected[0] else ops.add(corrected[i], delta) + } + } + + // 5 — merge streams, final norm, tied lm_head (gemma3n has no final softcap). + val merged = altupGlobals.mergeStreams(streams, ctx) + return lmHead.forward(outputNorm.forward(merged, ctx), ctx) + } + + /** `per_layer_inputs[..., layerIdx, :]` matched to the trunk's working rank. */ + private fun perLayerSlice( + perLayerInputs: Tensor, + layerIdx: Int, + workingRank: Int, + ctx: ExecutionContext, + ): Tensor { + val ops = ctx.ops + // [B, S, L, pleDim] → narrow L → [B, S, 1, pleDim] → squeeze → [B, S, pleDim] + val slice = ops.squeeze(ops.narrow(perLayerInputs, dim = 2, start = layerIdx, length = 1), dim = 2) + return if (workingRank == 2) ops.squeeze(slice, dim = 0) else slice + } +} diff --git a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nModelMetadata.kt b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nModelMetadata.kt index 85045df1..d3225dd8 100644 --- a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nModelMetadata.kt +++ b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nModelMetadata.kt @@ -34,8 +34,19 @@ public data class Gemma3nModelMetadata( /** Per-layer activation sparsity rates. Empty means no sparsity. */ val activationSparsityPattern: List = emptyList(), /** Activation sparsity scale factor (from GGUF: gemma3n.activation_sparsity_scale). */ - val activationSparsityScale: Float = 0f + val activationSparsityScale: Float = 0f, + /** RMSNorm epsilon (`gemma3n.attention.layer_norm_rms_epsilon`; real E2B/E4B: 1e-6). */ + val rmsNormEps: Float = 1e-6f, + /** + * Per-layer activation-sparsity std multipliers (`gemma3n.activation_sparsity_scale`). + * Real checkpoints store `Φ⁻¹(target_sparsity) ≈ 1.6449` on sparse layers and `-inf` + * on the rest — non-finite (or empty list) disables sparsity for that layer. + */ + val activationSparsityScales: List = emptyList(), ) { + /** The std multiplier for one layer, or `null` when sparsity is off there. */ + public fun sparsityScaleFor(layerIdx: Int): Float? = + activationSparsityScales.getOrNull(layerIdx)?.takeIf { it.isFinite() && it > 0f } /** * Returns the layer type at the given layer index. * Pattern repeats: ["sliding", "sliding", "sliding", "sliding", "full"] diff --git a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nNetworkDef.kt b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nNetworkDef.kt new file mode 100644 index 00000000..ee4b03fc --- /dev/null +++ b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nNetworkDef.kt @@ -0,0 +1,191 @@ +package sk.ainet.models.gemma3n + +import sk.ainet.apps.llm.HybridTransformerBlock +import sk.ainet.lang.nn.DefaultNeuralNetworkExecutionContext +import sk.ainet.lang.nn.Module +import sk.ainet.lang.nn.dsl.NeuralNetworkDslImpl +import sk.ainet.lang.nn.dsl.StageImpl +import sk.ainet.lang.nn.dsl.embedding +import sk.ainet.lang.nn.dsl.multiHeadAttention +import sk.ainet.lang.nn.dsl.rmsNorm +import sk.ainet.lang.nn.layers.EmbeddingAdapter +import sk.ainet.lang.nn.normalization.RMSNormalization +import sk.ainet.lang.nn.transformer.OwnerReadOnlyKVCache +import sk.ainet.lang.nn.transformer.PositionalKVCache +import sk.ainet.lang.nn.transformer.RoPEMode +import sk.ainet.lang.nn.transformer.VoidDense +import sk.ainet.lang.types.DType +import sk.ainet.models.gemma.LayerType +import sk.ainet.models.gemma.PerLayerEmbedding +import kotlin.reflect.KClass + +/** + * Gemma 3n architecture defined via the network DSL — the #377 DSL migration, replacing the + * hand-rolled `Gemma3nRuntime`. Faithful to HF `Gemma3nTextModel` (see [Gemma3nModel] for + * the per-layer flow) and buildable into a compute-graph tape for the StableHLO → IREE + * mobile path. + * + * What gemma3n adds over the gemma-4 lane's `gemmaNetwork()` (which already carries hybrid + * sliding/global attention, dual RoPE bases, per-layer FFN dims, per-type shared KV, PLE, + * q/k-norm, parameterless v-norm and attention scale 1.0): + * [Gemma3nAltUpBlock] (4 parallel streams + router), [Gemma3nLaurelBlock], + * [Gemma3nSparseGeGluFFN] (Gaussian-top-k on the first layers) and the PLE delta going to + * the non-active AltUp streams instead of the residual. + */ +public fun gemma3nNetwork( + metadata: Gemma3nModelMetadata, + dtype: KClass, + maxInferenceLen: Int = minOf(metadata.contextLength, 4096), + /** HF `laurel_rank` (64 on real checkpoints; not in the GGUF — the loader derives it + * from `blk.0.laurel_l`'s shape). */ + laurelRank: Int = LAUREL_RANK, +): Module { + val dim = metadata.embeddingLength + val nHeads = metadata.headCount + val nKVHeads = metadata.kvHeadCount + val nLayers = metadata.blockCount + val headDim = metadata.headDim + val seqLen = maxInferenceLen + val vocabSize = metadata.vocabSize + val eps = metadata.rmsNormEps + + val nnCtx = DefaultNeuralNetworkExecutionContext() + val dslImpl = NeuralNetworkDslImpl(nnCtx, dtype) + dslImpl.embedding(vocabSize, dim, id = "token_embd") + + // KV sharing: same owner-per-attention-type scheme as gemma-4 (HF: a shared layer + // reuses the K/V of the LAST non-shared layer of the same type). + val firstSharedLayer = nLayers - metadata.kvSharedLayers + val typeOwners = mutableMapOf>() + val typeOwnerLayerIdx = mutableMapOf() + if (metadata.kvSharedLayers > 0) { + for (l in 0 until firstSharedLayer) { + typeOwnerLayerIdx[metadata.getLayerType(l)] = l + } + } + + for (layer in 0 until nLayers) { + val layerType = metadata.getLayerType(layer) + val isGlobal = layerType == LayerType.GLOBAL + val ropeBase = metadata.getRopeBase(layer) + val slidingWindow = if (isGlobal) null else metadata.slidingWindow + val ffnDim = metadata.feedForwardLengths.getOrElse(layer) { metadata.feedForwardLengths.last() } + val isInSharedGroup = metadata.kvSharedLayers > 0 && layer >= firstSharedLayer + + val stage = StageImpl(nnCtx, "blk.$layer", dtype) + stage.rmsNorm(dim, eps, id = "attn_norm", unitOffset = false) + stage.multiHeadAttention( + dim = dim, + nHeads = nHeads, + nKVHeads = nKVHeads, + causal = true, + // HF Gemma3nTextAttention: per-head RMSNorm (with scale) on Q and K before RoPE, + // parameterless per-head RMSNorm on V, attention scaling fixed to 1.0. + qkNorm = true, + qkNormUnitOffset = false, + qkNormEps = eps, + attentionScale = 1.0f, + vNormNoScale = true, + id = "attn", + slidingWindow = slidingWindow, + ) { + rope( + headDim = headDim, + maxSeqLen = seqLen, + mode = RoPEMode.SPLIT_HALF, + base = ropeBase, + ) + if (!isInSharedGroup) { + val own = PositionalKVCache( + maxSeqLen = seqLen, + nKVHeads = nKVHeads, + headDim = headDim, + name = "blk.$layer.attn.kv_cache", + ) + kvCache(own) + if (typeOwnerLayerIdx[layerType] == layer) typeOwners[layerType] = own + } else { + val ownerCache = typeOwners[layerType] + ?: error( + "gemma3n: kv-shared layer $layer (type=$layerType) has no non-shared " + + "owner of the same type before firstSharedLayer=$firstSharedLayer", + ) + kvCache(OwnerReadOnlyKVCache(delegate = ownerCache, name = "blk.$layer.attn.kv_cache")) + } + } + stage.rmsNorm(dim, eps, id = "post_attention_norm", unitOffset = false) + stage.rmsNorm(dim, eps, id = "ffn_norm", unitOffset = false) // pre_feedforward_layernorm + stage.modules += Gemma3nSparseGeGluFFN( + hiddenSize = dim, + ffnDim = ffnDim, + stdMultiplier = metadata.sparsityScaleFor(layer) ?: Float.NEGATIVE_INFINITY, + dtype = dtype, + name = "ffn", + ) + stage.rmsNorm(dim, eps, id = "post_ffw_norm", unitOffset = false) + stage.modules += Gemma3nLaurelBlock( + hiddenSize = dim, + laurelRank = laurelRank, + rmsEps = eps, + dtype = dtype, + name = "laurel", + ) + stage.modules += Gemma3nAltUpBlock( + hiddenSize = dim, + numInputs = metadata.numAltupInputs, + activeIdx = metadata.altupActiveIdx, + rmsEps = eps, + dtype = dtype, + name = "altup", + ) + stage.modules += Gemma3nPerLayerApply( + hiddenSize = dim, + perLayerDim = metadata.perLayerEmbeddingLength, + rmsEps = eps, + dtype = dtype, + name = "per_layer_input", + ) + dslImpl.modules += HybridTransformerBlock(stage.modules.toList(), name = "blk.$layer") + } + + dslImpl.rmsNorm(dim, eps, id = "output_norm", unitOffset = false) + // Void placeholder — the 262k vocab head would eagerly allocate ~2 GB of zeros otherwise; + // gemma3n ties the head to token_embd, bound by the loader. + dslImpl.modules += VoidDense("output", vocabSize, dim, dtype = dtype) + + @Suppress("UNCHECKED_CAST") + val tokenEmbedding = dslImpl.modules[0] as EmbeddingAdapter + val blocks = dslImpl.modules.filterIsInstance>() + val outputNorm = dslImpl.modules[dslImpl.modules.size - 2] as RMSNormalization + @Suppress("UNCHECKED_CAST") + val lmHead = dslImpl.modules[dslImpl.modules.size - 1] as VoidDense + + val ple = PerLayerEmbedding( + vocabSize = vocabSize, + hiddenSize = dim, + numLayers = nLayers, + perLayerDim = metadata.perLayerEmbeddingLength, + rmsEps = eps, + ) + + return Gemma3nModel( + tokenEmbedding = tokenEmbedding, + ple = ple, + altupGlobals = Gemma3nAltUpGlobals( + hiddenSize = dim, + numInputs = metadata.numAltupInputs, + dtype = dtype, + ), + blocks = blocks, + outputNorm = outputNorm, + lmHead = lmHead, + dtype = dtype, + activeIdx = metadata.altupActiveIdx, + // HF: embed_scale = hidden_size ** 0.5 (bf16-rounded in the reference; we bisect + // against llama.cpp if the rounding ever matters at parity tolerance). + embedScale = kotlin.math.sqrt(dim.toFloat()), + ) +} + +/** HF `laurel_rank` — constant 64 across gemma-3n checkpoints (not stored in the GGUF). */ +public const val LAUREL_RANK: Int = 64 diff --git a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nNetworkLoader.kt b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nNetworkLoader.kt new file mode 100644 index 00000000..0af2179d --- /dev/null +++ b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nNetworkLoader.kt @@ -0,0 +1,84 @@ +package sk.ainet.models.gemma3n + +import sk.ainet.context.ExecutionContext +import sk.ainet.io.RandomAccessSource +import sk.ainet.io.weights.MappingConfig +import sk.ainet.io.weights.WeightMapper +import sk.ainet.io.weights.WeightTensor +import sk.ainet.lang.nn.Module +import sk.ainet.lang.types.DType +import kotlin.reflect.KClass + +/** + * End-to-end loader for the Gemma 3n DSL path (#377): loads a real gemma3n GGUF through the + * engine-delegated [Gemma3nWeightLoader] (packed/MAPPED by default; the PLE table stays + * packed with row-dequant), builds [gemma3nNetwork] and binds every weight via + * [WeightMapper] + [Gemma3nGGUFNameResolver]. + */ +public object Gemma3nNetworkLoader { + + public suspend inline fun fromGguf( + ctx: ExecutionContext, + noinline randomAccessProvider: () -> RandomAccessSource, + maxInferenceLen: Int? = null, + debug: Boolean = false, + ): Module { + val weights = Gemma3nWeightLoader(randomAccessProvider).loadToMapStreaming(ctx) + return fromWeights(ctx, weights, T::class, maxInferenceLen, debug) + } + + public fun fromWeights( + ctx: ExecutionContext, + weights: Gemma3nWeights, + dtype: KClass, + maxInferenceLen: Int? = null, + debug: Boolean = false, + ): Module { + val md = weights.metadata + // laurel_rank is not a GGUF field — read it off the checkpoint's own tensor. + val laurelRank = weights.tensors["blk.0.laurel_l.weight"]?.shape?.get(0) ?: LAUREL_RANK + + val model = gemma3nNetwork( + md, + dtype, + maxInferenceLen = maxInferenceLen ?: minOf(md.contextLength, 4096), + laurelRank = laurelRank, + ) + + val weightTensors = weights.tensors.map { (name, tensor) -> + WeightTensor(name = name, shape = tensor.shape.dimensions.toList(), tensor = tensor) + } + val config = MappingConfig( + usePathBasedMatching = false, + fallbackToShapeMatching = false, + debug = debug, + nameResolver = Gemma3nGGUFNameResolver(), + ) + val result = WeightMapper.applyWeights(model, weightTensors, config) + + // gemma3n has no bias tensors; every non-bias DSL param must bind, and no loaded + // tensor may silently go unused (the qwen-bias lesson, transformers#352). + val unmappedNonBias = result.missingParams.filter { !it.contains(".bias") } + require(unmappedNonBias.isEmpty()) { + buildString { + appendLine("gemma3n: failed to map ${unmappedNonBias.size} weight parameters:") + unmappedNonBias.take(20).forEach { appendLine(" - $it") } + if (result.unusedTensors.isNotEmpty()) { + appendLine("Unused tensors (${result.unusedTensors.size}):") + result.unusedTensors.take(20).forEach { appendLine(" - $it") } + } + }.trim() + } + require(result.unusedTensors.isEmpty()) { + "gemma3n: tensors present in the GGUF but never bound: ${result.unusedTensors.take(20)}" + } + return model + } + + public inline fun fromWeights( + ctx: ExecutionContext, + weights: Gemma3nWeights, + maxInferenceLen: Int? = null, + debug: Boolean = false, + ): Module = fromWeights(ctx, weights, T::class, maxInferenceLen, debug) +} diff --git a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nWeightLoader.kt b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nWeightLoader.kt index 315442a6..d65bf8a6 100644 --- a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nWeightLoader.kt +++ b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nWeightLoader.kt @@ -386,9 +386,14 @@ public class Gemma3nWeightLoader private constructor( // Gemma 3n specific val slidingWindow = fields["$prefix.attention.sliding_window"]?.toIntValue() ?: Gemma3nModelMetadata.DEFAULT_SLIDING_WINDOW + // Real llama.cpp gemma3n GGUFs declare only `rope.freq_base` (the global/full base, + // 1M) — the sliding/local base is the SWA default 10k (llama.cpp + // `rope_freq_base_train_swa`). Legacy keys tried first for synthetic fixtures. val ropeBaseLocal = fields["$prefix.rope.freq_base_local"]?.toFloatValue() + ?: fields["$prefix.rope.freq_base_swa"]?.toFloatValue() ?: Gemma3nModelMetadata.DEFAULT_ROPE_BASE_LOCAL val ropeBaseGlobal = fields["$prefix.rope.freq_base_global"]?.toFloatValue() + ?: fields["$prefix.rope.freq_base"]?.toFloatValue() ?: Gemma3nModelMetadata.DEFAULT_ROPE_BASE_GLOBAL val kvSharedLayers = fields["$prefix.kv_shared_layers"]?.toIntValue() ?: fields["$prefix.attention.shared_kv_layers"]?.toIntValue() @@ -408,6 +413,15 @@ public class Gemma3nWeightLoader private constructor( // Activation sparsity val activationSparsityPattern = extractStreamingActivationSparsityPattern(fields, prefix) val activationSparsityScale = fields["$prefix.activation_sparsity_scale"]?.toFloatValue() ?: 0f + // Real GGUFs store the per-layer std multipliers directly (1.6449 on sparse layers, + // -inf on the rest); the DSL path consumes this list. + val activationSparsityScales = when (val v = fields["$prefix.activation_sparsity_scale"]) { + is List<*> -> v.mapNotNull { (it as? Number)?.toFloat() } + is FloatArray -> v.toList() + is DoubleArray -> v.map { it.toFloat() } + else -> emptyList() + } + val rmsNormEps = fields["$prefix.attention.layer_norm_rms_epsilon"]?.toFloatValue() ?: 1e-6f return Gemma3nModelMetadata( architecture = arch, @@ -428,7 +442,9 @@ public class Gemma3nWeightLoader private constructor( numAltupInputs = numAltupInputs, altupActiveIdx = altupActiveIdx, activationSparsityPattern = activationSparsityPattern, - activationSparsityScale = activationSparsityScale + activationSparsityScale = activationSparsityScale, + rmsNormEps = rmsNormEps, + activationSparsityScales = activationSparsityScales, ) } @@ -483,6 +499,11 @@ public class Gemma3nWeightLoader private constructor( if (perLayerValue != null && perLayerValue is List<*>) { return perLayerValue.mapNotNull { (it as? Number)?.toInt() } } + // Real llama.cpp GGUFs store the per-layer array under the SINGULAR key. + val singular = fields["$prefix.feed_forward_length"] + if (singular is List<*> && singular.size > 1) { + return singular.mapNotNull { (it as? Number)?.toInt() } + } // Fall back to single FFN length val ffnLength = fields["$prefix.feed_forward_length"]?.toIntValue() ?: (embeddingLength * 4) @@ -508,6 +529,11 @@ public class Gemma3nWeightLoader private constructor( if (patternValue != null && patternValue is List<*>) { return patternValue.mapNotNull { it as? String } } + // Real llama.cpp GGUFs store per-layer booleans: true = sliding, false = full. + val swaPattern = fields["$prefix.attention.sliding_window_pattern"] + if (swaPattern is List<*> && swaPattern.isNotEmpty() && swaPattern.first() is Boolean) { + return swaPattern.map { if (it == true) "sliding" else "full" } + } return Gemma3nModelMetadata.DEFAULT_LAYER_PATTERN } diff --git a/llm-inference/gemma3n/src/jvmTest/kotlin/sk/ainet/models/gemma3n/Gemma3nGoldenTokenParityTest.kt b/llm-inference/gemma3n/src/jvmTest/kotlin/sk/ainet/models/gemma3n/Gemma3nGoldenTokenParityTest.kt new file mode 100644 index 00000000..bf992f9f --- /dev/null +++ b/llm-inference/gemma3n/src/jvmTest/kotlin/sk/ainet/models/gemma3n/Gemma3nGoldenTokenParityTest.kt @@ -0,0 +1,100 @@ +package sk.ainet.models.gemma3n + +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlinx.coroutines.runBlocking +import sk.ainet.apps.llm.OptimizedLLMMode +import sk.ainet.apps.llm.OptimizedLLMRuntime +import sk.ainet.apps.llm.sampleFromTensor +import sk.ainet.apps.llm.tokenizer.TokenizerFactory +import sk.ainet.context.DirectCpuExecutionContext +import sk.ainet.io.JvmRandomAccessSource +import sk.ainet.io.gguf.StreamingGGUFReader +import sk.ainet.lang.types.FP32 + +/** + * The #346 maturity-gate parity probe for the Gemma 3n family (#377 DSL migration), against + * **mainline llama.cpp**: full 32-step greedy text equality on the DSL path + * ([Gemma3nWeightLoader] engine loading, packed/MAPPED → [Gemma3nNetworkLoader.fromWeights] + * → [OptimizedLLMRuntime]). This is the first time the family has a reference gate at all — + * and it exercises everything gemma3n adds over gemma-4: AltUp's four parallel streams with + * the tanh router, Laurel, Gaussian-top-k activation sparsity on the first ten layers, PLE + * feeding the non-active streams, and per-type shared KV for the last ten layers. + * + * Model-gated: runs only when `GEMMA3N_E2B_GGUF` points at `gemma-3n-E2B-it-Q4_K_M.gguf` + * AND the test JVM has ≥ 16 GB heap (`-PgemmaTestMaxHeap=20g`); skips quietly otherwise. + * The fixture header records the exact oracle build and commands. + */ +@org.junit.jupiter.api.Tag("smoke-reference") +@org.junit.jupiter.api.Tag("integration") +class Gemma3nGoldenTokenParityTest { + + private data class Fixture( + val prompt: String, + val steps: Int, + val promptTokens: List, + val oracleText: String, + ) + + private fun loadFixture(): Fixture { + val raw = checkNotNull(javaClass.getResourceAsStream("/gemma3n-e2b/golden-greedy-e2b.txt")) { + "fixture /gemma3n-e2b/golden-greedy-e2b.txt missing from test resources" + }.bufferedReader().readLines() + val map = raw.filter { it.isNotBlank() && !it.startsWith("#") } + .associate { it.substringBefore('=') to it.substringAfter('=') } + return Fixture( + prompt = map.getValue("prompt"), + steps = map.getValue("steps").toInt(), + promptTokens = map.getValue("prompt_tokens").split(',').map { it.trim().toInt() }, + oracleText = map.getValue("oracle_text"), + ) + } + + @Test + fun greedyDecodeMatchesMainlineLlamaCpp() { + val modelPath = System.getenv("GEMMA3N_E2B_GGUF") + if (modelPath.isNullOrBlank()) { + println("PARITY skipped: GEMMA3N_E2B_GGUF not set") + return + } + val maxHeapGb = Runtime.getRuntime().maxMemory() / (1024L * 1024L * 1024L) + if (maxHeapGb < 16) { + println("PARITY skipped: heap=$maxHeapGb GB < 16 GB; rerun with -PgemmaTestMaxHeap=20g") + return + } + val fixture = loadFixture() + val ctx = DirectCpuExecutionContext() + + // 1 — prompt tokenization parity (a failure here names the tokenizer, not the model). + val fields = StreamingGGUFReader.open(JvmRandomAccessSource.open(modelPath)).use { it.fields } + val tokenizer = TokenizerFactory.fromGgufFields(fields) + // gemma adds a BOS token (id 2); mirror llama.cpp's add_special encoding. + val raw = tokenizer.encode(fixture.prompt) + val encoded = if (raw.isNotEmpty() && raw[0] == tokenizer.bosTokenId) raw + else intArrayOf(tokenizer.bosTokenId) + raw + assertEquals( + fixture.promptTokens, encoded.toList(), + "prompt tokenization must match mainline llama.cpp", + ) + + // 2 — greedy continuation on the DSL path. + val weights = runBlocking { + Gemma3nWeightLoader( + randomAccessProvider = { JvmRandomAccessSource.open(modelPath) }, + ).loadToMapStreaming(ctx) + } + val model = Gemma3nNetworkLoader.fromWeights(ctx, weights) + val runtime = OptimizedLLMRuntime(model, ctx, OptimizedLLMMode.DIRECT, FP32::class) + for (i in 0 until fixture.promptTokens.size - 1) runtime.forward(fixture.promptTokens[i]) + var token = fixture.promptTokens.last() + val text = StringBuilder() + repeat(fixture.steps) { + token = sampleFromTensor(runtime.forward(token), 0f) + text.append(tokenizer.decode(token)) + } + assertEquals( + fixture.oracleText, text.toString(), + "greedy decode must equal mainline llama.cpp token-for-token", + ) + } +} diff --git a/llm-inference/gemma3n/src/jvmTest/resources/gemma3n-e2b/golden-greedy-e2b.txt b/llm-inference/gemma3n/src/jvmTest/resources/gemma3n-e2b/golden-greedy-e2b.txt new file mode 100644 index 00000000..955523b0 --- /dev/null +++ b/llm-inference/gemma3n/src/jvmTest/resources/gemma3n-e2b/golden-greedy-e2b.txt @@ -0,0 +1,19 @@ +# gemma-3n-E2B-it Q4_K_M greedy parity fixture (#377 DSL migration) +# +# Oracle: MAINLINE llama.cpp (brew, version 0.3.0 build 10621 commit c1d0e7a00). Asserts FULL +# cross-implementation greedy text equality on the new gemma3n DSL path +# (Gemma3nWeightLoader engine loading, packed/MAPPED -> gemma3nNetwork() -> +# OptimizedLLMRuntime) — AltUp (4 streams + router), Laurel, first-10-layer activation +# sparsity, PLE into the non-active streams, per-type shared KV (last 10 layers), hybrid +# sliding/global attention with dual RoPE bases: all exercised end-to-end. +# The prompt is chosen for greedy decisiveness (min top-1/top-2 gap 1.77 nats across all 32 +# steps; "The capital of France is" hits a 0.12-nat tie at step 3 that Q4_K +# cross-implementation noise legitimately flips). NOTE: gemma adds a BOS token (id 2) — +# prompt_tokens is the oracle's add_special encoding, BOS included. +# llama-server -m gemma-3n-E2B-it-Q4_K_M.gguf -ngl 0 -t 4 --port 8816 +# curl :8816/tokenize -d '{"content": "The first ten prime numbers are", "add_special": true}' +# curl :8816/completion -d '{"prompt": "The first ten prime numbers are", "n_predict": 32, "temperature": 0, "top_k": 1}' +prompt=The first ten prime numbers are +steps=32 +prompt_tokens=2,818,1171,3595,8355,4945,659 +oracle_text= 2, 3, 5, 7, 11, 13, 17, 19, 23, diff --git a/llm-runtime/kgemma/src/jvmMain/kotlin/sk/ainet/apps/kgemma/cli/Main.kt b/llm-runtime/kgemma/src/jvmMain/kotlin/sk/ainet/apps/kgemma/cli/Main.kt index 874e4dde..e7fbcfab 100644 --- a/llm-runtime/kgemma/src/jvmMain/kotlin/sk/ainet/apps/kgemma/cli/Main.kt +++ b/llm-runtime/kgemma/src/jvmMain/kotlin/sk/ainet/apps/kgemma/cli/Main.kt @@ -234,19 +234,32 @@ fun main(args: Array) { } } GemmaVariant.GEMMA3N -> { - val ingestion = Gemma3nIngestion( - ctx = ctx, - dtype = FP32::class, - config = Gemma3nLoadConfig() - ) when (format) { ModelFormat.GGUF -> { - println("Loading Gemma 3n GGUF model from $modelPath (streaming mode)...") - ingestion.loadRuntimeStreaming { - JvmRandomAccessSource.open(modelPath.toString()) + // The DSL lane (#377): engine loading (packed/MAPPED), + // gemma3nNetwork() with AltUp/Laurel/sparsity/PLE, verified + // token-for-token vs mainline llama.cpp by + // Gemma3nGoldenTokenParityTest. The hand-rolled + // Gemma3nRuntime never applied PLE and predates the gate. + println("Loading Gemma 3n GGUF model from $modelPath via gemma3nNetwork() + OptimizedLLMRuntime (engine loader, keep-packed, mapped)...") + val model = kotlinx.coroutines.runBlocking { + sk.ainet.models.gemma3n.Gemma3nNetworkLoader.fromGguf( + ctx, + { JvmRandomAccessSource.open(modelPath.toString()) }, + ) } + sk.ainet.apps.llm.OptimizedLLMRuntime( + model, ctx, sk.ainet.apps.llm.OptimizedLLMMode.DIRECT, FP32::class, + ) } ModelFormat.SAFETENSORS -> { + // SafeTensors stays on the legacy hand-rolled runtime until the DSL + // lane grows a SafeTensors leg (tracked with the #377 remainder). + val ingestion = Gemma3nIngestion( + ctx = ctx, + dtype = FP32::class, + config = Gemma3nLoadConfig() + ) val modelDir = if (modelPath.isDirectory()) modelPath else modelPath.parent ?: modelPath val indexPath = modelDir.resolve("model.safetensors.index.json") val safetensorsPath = if (indexPath.exists()) indexPath.toString() diff --git a/tests/smoke/smoke-models.json b/tests/smoke/smoke-models.json index bed73a85..9d578431 100644 --- a/tests/smoke/smoke-models.json +++ b/tests/smoke/smoke-models.json @@ -63,6 +63,14 @@ "format": "gguf", "steps": 16 }, + { + "name": "Gemma3n-E2B-GGUF", + "runner": "skainet", + "model": "gemma-3n-E2B-it-Q4_K_M.gguf", + "format": "gguf", + "steps": 12, + "prompt": "The first ten prime numbers are" + }, { "name": "Gemma4-E4B-GGUF", "runner": "kgemma",