From 192b903c3504a2bc7161cdd0d784e27c8057eb76 Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Wed, 2 Sep 2026 13:14:58 +0200 Subject: [PATCH 1/3] refactor(gemma): SafeTensors loading collapses onto the engine's ShardedSafeTensorsParametersLoader (SKaiNET#1246) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit GemmaSafeTensorsLoader hand-rolled the per-tensor materialization the engine's single-file loader already owned — BF16/F16 widening via the GGUF module's DequantOps (a wrong dependency edge), a row-major transpose no call site ever enabled, and a PLE size guard — because the engine had no sharded ParametersLoader. SKaiNET 0.53.0 ships one (SKaiNET#1252). The loader now expresses only family policy: the HF allowlist and PLE size guard as the engine loader's tensorFilter, and the HF -> GGUF slot renaming (table-driven per layer). Every dtype decision is the engine's, driven by a DTypePolicy forwarded through ShardedSafeTensorsParametersLoader.withPolicy; GemmaNetworkLoader's SafeTensors lane validates against the same keep-native set as the GGUF lane instead of rejecting every Require(). transposeRowMajor and the DequantOps import are gone; GemmaSafeTensorsMappedPle and the PLE auto-disable semantics are unchanged. Co-Authored-By: Claude Fable 5.1 --- .../ainet/models/gemma/GemmaNetworkLoader.kt | 14 +- .../models/gemma/GemmaSafeTensorsLoader.kt | 325 +++++++----------- 2 files changed, 123 insertions(+), 216 deletions(-) diff --git a/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaNetworkLoader.kt b/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaNetworkLoader.kt index c8c03ba2..14d79af9 100644 --- a/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaNetworkLoader.kt +++ b/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaNetworkLoader.kt @@ -45,8 +45,8 @@ public class GemmaNetworkLoader @PublishedApi internal constructor( // The GGUF lane honors narrow-float keep-native exactly as llama's shared loader does // (GemmaWeightLoader resolves keepF16Native/keepBf16Native from the policy) — the old // `keepNative = emptySet()` claim predated that and silently rejected `Require(BF16)` - // on a loader that supports it (#375). The SafeTensors lane still widens; it re-validates - // against an empty set at its own dispatch site below. + // on a loader that supports it (#375). The SafeTensors lane now rides the engine's + // sharded loader (SKaiNET#1246) and keeps native under the same set. DTypePolicyValidation.validate( policy, "GemmaNetworkLoader.withDtypePolicy", keepNative = sk.ainet.lang.nn.dsl.decoder.DECODER_NARROW_KEEP_NATIVE, @@ -146,12 +146,14 @@ public class GemmaNetworkLoader @PublishedApi internal constructor( loader.loadToMapStreaming(ctx) } is WeightsProvider.SafeTensorsIndex -> { - // The hand-rolled SafeTensors reader widens every narrow float (#375): reject a - // Require() policy here rather than accept-and-ignore it. + // The SafeTensors lane rides the engine's ShardedSafeTensorsParametersLoader + // (SKaiNET#1246), which keeps BF16/FP16 native under Require() exactly as the + // GGUF lane does — the same keep-native set applies. DTypePolicyValidation.validate( - dtypePolicy, "GemmaNetworkLoader(SafeTensors)", keepNative = emptySet(), + dtypePolicy, "GemmaNetworkLoader(SafeTensors)", + keepNative = sk.ainet.lang.nn.dsl.decoder.DECODER_NARROW_KEEP_NATIVE, ) - val loader = GemmaSafeTensorsLoader(wp.indexPath) + val loader = GemmaSafeTensorsLoader(wp.indexPath, dtypePolicy = dtypePolicy) @Suppress("UNCHECKED_CAST") loader.loadToMap(ctx, T::class) as GemmaWeights } diff --git a/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaSafeTensorsLoader.kt b/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaSafeTensorsLoader.kt index a26f57bf..8fc45732 100644 --- a/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaSafeTensorsLoader.kt +++ b/llm-inference/gemma/src/commonMain/kotlin/sk/ainet/models/gemma/GemmaSafeTensorsLoader.kt @@ -1,13 +1,13 @@ package sk.ainet.models.gemma import sk.ainet.context.ExecutionContext -import sk.ainet.io.gguf.dequant.DequantOps +import sk.ainet.io.load +import sk.ainet.io.safetensors.ShardedSafeTensorsParametersLoader import sk.ainet.io.safetensors.ShardedTensorInfo -import sk.ainet.io.safetensors.StreamingShardedSafeTensorsReader import sk.ainet.io.safetensors.readTextFile -import sk.ainet.lang.tensor.Shape import sk.ainet.lang.tensor.Tensor import sk.ainet.lang.types.DType +import sk.ainet.lang.types.DTypePolicy import sk.ainet.lang.types.FP32 import kotlin.reflect.KClass @@ -21,15 +21,22 @@ import kotlin.reflect.KClass * - Uses GemmaConfigParser instead of Gemma3nConfigParser * - No AltUp, Laurel, or activation sparsity tensors * - Per-layer head dim may vary (global_head_dim vs head_dim) - * - * Engine-delegation status (#375): the sharded READING already rides the engine - * (`StreamingShardedSafeTensorsReader` / `SafeTensorsIndexParser`); what remains family-side is - * the per-tensor materialization policy (bf16/f16 widening, row-major transpose, size guards), - * because the engine's `SafeTensorsParametersLoader` is single-file only. Once SKaiNET#1246 - * ships a sharded ParametersLoader with narrow-float policies, this class collapses onto it. + * + * Engine delegation (#375 → SKaiNET#1246): reading AND per-tensor materialization ride the + * engine's [ShardedSafeTensorsParametersLoader]. This class owns only the family policy — + * which HF tensors are wanted (the `tensorFilter`), the HF → GGUF slot renaming, and the PLE + * size guard — while every dtype decision (BF16/F16 widening vs keep-native, F32 passthrough) + * is the engine's, driven by [dtypePolicy] exactly as on the GGUF lane. No per-family dequant + * code remains here (#346). + * + * @param indexPath path to `model.safetensors.index.json`. + * @param dtypePolicy narrow-float policy forwarded to + * [ShardedSafeTensorsParametersLoader.withPolicy]: `Any` (default) widens BF16/F16 to FP32; + * `Require(BF16)` / `Require(FP16)` keep the native encoding under an FP32-typed tensor. */ public class GemmaSafeTensorsLoader( - private val indexPath: String + private val indexPath: String, + private val dtypePolicy: DTypePolicy = DTypePolicy.Any, ) { public suspend fun loadToMap( @@ -47,25 +54,25 @@ public class GemmaSafeTensorsLoader( val metadata = GemmaConfigParser.parseFromJson(configJson) - val reader = StreamingShardedSafeTensorsReader.openFromIndex(indexPath) - - return reader.use { - loadFromReader(ctx, dtype, it, metadata) - } - } + // Family policy, expressed as the engine loader's filter: only the tensors this + // family maps (vision/audio towers and anything unmapped never materialize, and are + // exempt from the fail-fast dtype pre-scan), and the PLE table only when it fits the + // eager ByteArray path — see the note in loadGlobalTensors. + val wanted = wantedHfNames(metadata.blockCount) + val loader = ShardedSafeTensorsParametersLoader.withPolicy( + indexPath = indexPath, + policy = dtypePolicy, + tensorFilter = { info: ShardedTensorInfo -> + info.name in wanted && (info.name !in PLE_TABLE_CANDIDATES || info.sizeInBytes <= MAX_BYTES_PER_TENSOR) + }, + ) + val tensorsByHfName = linkedMapOf>() + loader.load(ctx, dtype) { name, tensor -> tensorsByHfName[name] = tensor } - private fun loadFromReader( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - metadata: GemmaModelMetadata - ): GemmaWeights { val tensorsByGgufName = linkedMapOf>() - val tensorsByHfName = reader.tensors.associateBy { it.name } - - loadGlobalTensors(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName) + loadGlobalTensors(tensorsByHfName, tensorsByGgufName) for (layer in 0 until metadata.blockCount) { - loadLayerTensors(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, layer) + loadLayerTensors(tensorsByHfName, tensorsByGgufName, layer) } return GemmaWeights( @@ -75,23 +82,17 @@ public class GemmaSafeTensorsLoader( } private fun loadGlobalTensors( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, + tensorsByHfName: Map>, tensorsByGgufName: MutableMap> ) { // Token embeddings - val embedTokens = tensorsByHfName[HF_EMBED_TOKENS] + val embedTensor = tensorsByHfName[HF_EMBED_TOKENS] ?: error("Missing tensor: $HF_EMBED_TOKENS") - val embedTensor = loadAndConvertTensor(ctx, dtype, reader, embedTokens, transpose = false) tensorsByGgufName[GemmaTensorNames.TOKEN_EMBEDDINGS] = embedTensor // Output norm - val norm = tensorsByHfName[HF_OUTPUT_NORM] + tensorsByGgufName[GemmaTensorNames.OUTPUT_NORM] = tensorsByHfName[HF_OUTPUT_NORM] ?: error("Missing tensor: $HF_OUTPUT_NORM") - val normTensor = loadAndConvertTensor(ctx, dtype, reader, norm, transpose = false) - tensorsByGgufName[GemmaTensorNames.OUTPUT_NORM] = normTensor // Output weight (weight tying - reuse embed_tokens) tensorsByGgufName[GemmaTensorNames.OUTPUT_WEIGHT] = embedTensor @@ -103,207 +104,111 @@ public class GemmaSafeTensorsLoader( // // The token-embeddings-per-layer tensor on real Gemma-4 E2B is // [vocab_size, num_layers, per_layer_dim] in BF16 — 4.7 GB raw, well - // over the 2 GB JVM ByteArray limit our reader uses. The GGUF path - // sidesteps this by keeping the table Q6_K-packed (~1.8 GB) with a - // dedicated row-dequant data type. We don't have an equivalent for - // SafeTensors yet, so when the BF16 source is too large we skip the - // load — leaving `per_layer_token_embd.weight` absent from the - // tensor map, which auto-disables PLE in `GemmaNetworkLoader`. The - // model still runs (sandwich norms, layer_output_scale and softcap - // are intact) but without the per-layer side-channel signal. PLE - // support on the SafeTensors path is tracked separately. - loadFirstExistingIfFits( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - ggufName = GemmaTensorNames.PER_LAYER_TOKEN_EMBD, - hfCandidates = listOf(HF_EMBED_TOKENS_PER_LAYER, HF_PER_LAYER_TOKEN_EMBD), - transpose = false, - ) - loadFirstExisting( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - ggufName = GemmaTensorNames.PER_LAYER_MODEL_PROJ, - hfCandidates = listOf(HF_PER_LAYER_MODEL_PROJECTION, HF_PER_LAYER_MODEL_PROJ), - transpose = false, - ) - loadFirstExisting( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - ggufName = GemmaTensorNames.PER_LAYER_PROJ_NORM, - hfCandidates = listOf(HF_PER_LAYER_PROJECTION_NORM, HF_PER_LAYER_PROJ_NORM), - transpose = false, - ) + // over the 2 GB JVM ByteArray limit the eager reader path is bound by. + // The GGUF path sidesteps this by keeping the table Q6_K-packed + // (~1.8 GB) with a dedicated row-dequant data type. On SafeTensors the + // tensorFilter (see loadToMap) skips the table when it is too large — + // leaving `per_layer_token_embd.weight` absent from the tensor map, + // which auto-disables PLE in `GemmaNetworkLoader`. The model still + // runs (sandwich norms, layer_output_scale and softcap are intact) + // but without the per-layer side-channel signal; the JVM + // `GemmaSafeTensorsMappedPle` path injects it file-backed afterwards. + mapFirstExisting(tensorsByHfName, tensorsByGgufName, + GemmaTensorNames.PER_LAYER_TOKEN_EMBD, PLE_TABLE_CANDIDATES) + mapFirstExisting(tensorsByHfName, tensorsByGgufName, + GemmaTensorNames.PER_LAYER_MODEL_PROJ, listOf(HF_PER_LAYER_MODEL_PROJECTION, HF_PER_LAYER_MODEL_PROJ)) + mapFirstExisting(tensorsByHfName, tensorsByGgufName, + GemmaTensorNames.PER_LAYER_PROJ_NORM, listOf(HF_PER_LAYER_PROJECTION_NORM, HF_PER_LAYER_PROJ_NORM)) } private fun loadLayerTensors( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, + tensorsByHfName: Map>, tensorsByGgufName: MutableMap>, layer: Int ) { - // Pre-attention input layernorm (HF: input_layernorm; GGUF slot: attn_norm). - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfInputLayernorm(layer), GemmaTensorNames.inputLayernorm(layer), transpose = false) - - // Attention projections. - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnQ(layer), GemmaTensorNames.attnQ(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnK(layer), GemmaTensorNames.attnK(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnV(layer), GemmaTensorNames.attnV(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnO(layer), GemmaTensorNames.attnOut(layer)) - - // Sandwich norms — Gemma 4 has FOUR norms per block: - // - attn_norm ← input_layernorm (already loaded above) - // - post_attention_norm ← post_attention_layernorm (Gemma-4 only) - // - ffn_norm ← pre_feedforward_layernorm (Gemma-4 only) - // - post_ffw_norm ← post_feedforward_layernorm (Gemma-4 only) - // The GGUF naming on `GemmaTensorNames.postAttentionLayernorm` is a - // legacy alias for the pre-FFN norm slot (`ffn_norm`); the proper - // Gemma-4 source for that slot is HF `pre_feedforward_layernorm`. - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPostAttnLayernorm(layer), GemmaTensorNames.postAttentionNorm(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPreFfwLayernorm(layer), GemmaTensorNames.postAttentionLayernorm(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPostFfwLayernorm(layer), GemmaTensorNames.postFfwNorm(layer), transpose = false) - - // MLP weights. - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfMlpGate(layer), GemmaTensorNames.ffnGate(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfMlpUp(layer), GemmaTensorNames.ffnUp(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfMlpDown(layer), GemmaTensorNames.ffnDown(layer)) - - // Per-layer scalar (HF: `layer_scalar`, no `.weight` suffix; scalar shape). - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfLayerScalar(layer), GemmaTensorNames.layerOutputScale(layer), transpose = false) - - // Optional Per-Layer Embedding (PLE) per-layer tensors. - // The HF projection sends [B,S,perLayerDim] back into the residual at - // hidden_size, so it goes into the `proj` slot, NOT `per_layer_input`. - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPerLayerInputGate(layer), GemmaTensorNames.pleInpGate(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPerLayerProjection(layer), GemmaTensorNames.pleProj(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPostPerLayerInputNorm(layer), GemmaTensorNames.plePostNorm(layer), transpose = false) - - // Optional QK normalization (per-head Q/K RMSNorm). - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnQNorm(layer), GemmaTensorNames.attnQNorm(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnKNorm(layer), GemmaTensorNames.attnKNorm(layer), transpose = false) - } - - private fun loadTensorIfExists( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, - tensorsByGgufName: MutableMap>, - hfName: String, - ggufName: String, - transpose: Boolean = false - ) { - val info = tensorsByHfName[hfName] - if (info != null) { - val tensor = loadAndConvertTensor(ctx, dtype, reader, info, transpose) - tensorsByGgufName[ggufName] = tensor + for ((hfName, ggufName) in layerSlots(layer)) { + mapIfExists(tensorsByHfName, tensorsByGgufName, hfName, ggufName) } } - private fun loadFirstExisting( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, + private fun mapIfExists( + tensorsByHfName: Map>, tensorsByGgufName: MutableMap>, - ggufName: String, - hfCandidates: List, - transpose: Boolean = false + hfName: String, + ggufName: String ) { - for (hfName in hfCandidates) { - val info = tensorsByHfName[hfName] ?: continue - tensorsByGgufName[ggufName] = loadAndConvertTensor(ctx, dtype, reader, info, transpose) - return - } + tensorsByHfName[hfName]?.let { tensorsByGgufName[ggufName] = it } } - /** - * Variant of [loadFirstExisting] that silently skips the load if the - * raw on-disk tensor exceeds [MAX_BYTES_PER_TENSOR] (the JVM ByteArray - * limit our streaming reader is bound by). Used for PLE globals on - * real Gemma-4 SafeTensors checkpoints, where the BF16 source is too - * large for the eager reader path. - */ - private fun loadFirstExistingIfFits( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, + private fun mapFirstExisting( + tensorsByHfName: Map>, tensorsByGgufName: MutableMap>, ggufName: String, - hfCandidates: List, - transpose: Boolean = false + hfCandidates: List ) { for (hfName in hfCandidates) { - val info = tensorsByHfName[hfName] ?: continue - if (info.sizeInBytes > MAX_BYTES_PER_TENSOR) { - // Leave the slot empty so PLE auto-detection turns off cleanly. - return - } - tensorsByGgufName[ggufName] = loadAndConvertTensor(ctx, dtype, reader, info, transpose) + val tensor = tensorsByHfName[hfName] ?: continue + tensorsByGgufName[ggufName] = tensor return } } - @Suppress("UNCHECKED_CAST") - private fun loadAndConvertTensor( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - info: ShardedTensorInfo, - transpose: Boolean = false - ): Tensor { - val bytes = reader.loadTensorData(info) - val shape = Shape(*info.shape.map { it.toInt() }.toIntArray()) - - val floats = when (info.dtype.uppercase()) { - "BF16" -> DequantOps.dequantBF16FromBytes(bytes) - "F16" -> DequantOps.dequantF16FromBytes(bytes) - "F32" -> DequantOps.bytesToFloatArray(bytes) - else -> error("Unsupported SafeTensors dtype: ${info.dtype}") - } + private companion object { + // The engine's eager loader path returns each tensor as a single + // ByteArray, which the JVM caps at Int.MAX_VALUE bytes. Keep a + // little headroom for safety. + const val MAX_BYTES_PER_TENSOR: Long = Int.MAX_VALUE.toLong() - 1024L - return if (transpose && shape.rank == 2) { - val rows = shape[0] - val cols = shape[1] - val transposed = transposeRowMajor(floats, rows, cols) - val newShape = Shape(cols, rows) - ctx.fromFloatArray(newShape, dtype, transposed) as Tensor - } else { - ctx.fromFloatArray(shape, dtype, floats) as Tensor + /** HF names of the PLE table, new-name first; the only tensors subject to the size guard. */ + val PLE_TABLE_CANDIDATES: List by lazy { + listOf(HF_EMBED_TOKENS_PER_LAYER, HF_PER_LAYER_TOKEN_EMBD) } - } - private fun transposeRowMajor(data: FloatArray, rows: Int, cols: Int): FloatArray { - val out = FloatArray(data.size) - for (r in 0 until rows) { - for (c in 0 until cols) { - out[c * rows + r] = data[r * cols + c] - } + /** Every HF tensor this family maps, for [blockCount] layers — the engine loader's allowlist. */ + fun wantedHfNames(blockCount: Int): Set = buildSet { + add(HF_EMBED_TOKENS); add(HF_OUTPUT_NORM) + addAll(PLE_TABLE_CANDIDATES) + add(HF_PER_LAYER_MODEL_PROJECTION); add(HF_PER_LAYER_MODEL_PROJ) + add(HF_PER_LAYER_PROJECTION_NORM); add(HF_PER_LAYER_PROJ_NORM) + for (layer in 0 until blockCount) addAll(layerSlots(layer).map { it.first }) } - return out - } - private companion object { - // The streaming sharded reader returns each tensor as a single - // ByteArray, which the JVM caps at Int.MAX_VALUE bytes. Keep a - // little headroom for safety. - const val MAX_BYTES_PER_TENSOR: Long = Int.MAX_VALUE.toLong() - 1024L + /** + * Per-layer HF → GGUF slot pairs, in the order the GGUF lane populates them. + * + * Sandwich norms — Gemma 4 has FOUR norms per block: + * - attn_norm ← input_layernorm + * - post_attention_norm ← post_attention_layernorm (Gemma-4 only) + * - ffn_norm ← pre_feedforward_layernorm (Gemma-4 only) + * - post_ffw_norm ← post_feedforward_layernorm (Gemma-4 only) + * The GGUF naming on `GemmaTensorNames.postAttentionLayernorm` is a legacy alias for + * the pre-FFN norm slot (`ffn_norm`); the proper Gemma-4 source for that slot is HF + * `pre_feedforward_layernorm`. + * + * PLE: the HF projection sends [B,S,perLayerDim] back into the residual at + * hidden_size, so it goes into the `proj` slot, NOT `per_layer_input`. + */ + fun layerSlots(layer: Int): List> = listOf( + hfInputLayernorm(layer) to GemmaTensorNames.inputLayernorm(layer), + hfAttnQ(layer) to GemmaTensorNames.attnQ(layer), + hfAttnK(layer) to GemmaTensorNames.attnK(layer), + hfAttnV(layer) to GemmaTensorNames.attnV(layer), + hfAttnO(layer) to GemmaTensorNames.attnOut(layer), + hfPostAttnLayernorm(layer) to GemmaTensorNames.postAttentionNorm(layer), + hfPreFfwLayernorm(layer) to GemmaTensorNames.postAttentionLayernorm(layer), + hfPostFfwLayernorm(layer) to GemmaTensorNames.postFfwNorm(layer), + hfMlpGate(layer) to GemmaTensorNames.ffnGate(layer), + hfMlpUp(layer) to GemmaTensorNames.ffnUp(layer), + hfMlpDown(layer) to GemmaTensorNames.ffnDown(layer), + // Per-layer scalar (HF: `layer_scalar`, no `.weight` suffix; scalar shape). + hfLayerScalar(layer) to GemmaTensorNames.layerOutputScale(layer), + hfPerLayerInputGate(layer) to GemmaTensorNames.pleInpGate(layer), + hfPerLayerProjection(layer) to GemmaTensorNames.pleProj(layer), + hfPostPerLayerInputNorm(layer) to GemmaTensorNames.plePostNorm(layer), + // Optional QK normalization (per-head Q/K RMSNorm). + hfAttnQNorm(layer) to GemmaTensorNames.attnQNorm(layer), + hfAttnKNorm(layer) to GemmaTensorNames.attnKNorm(layer), + ) const val HF_EMBED_TOKENS = "model.language_model.embed_tokens.weight" const val HF_OUTPUT_NORM = "model.language_model.norm.weight" From e9593b1bcef4a65ffed19dd6771273251929773b Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Wed, 2 Sep 2026 13:17:51 +0200 Subject: [PATCH 2/3] test(gemma): synthetic sharded SafeTensors fixture for the engine-backed loader (SKaiNET#1246) A 2-shard checkpoint written with the engine's SafeTensorsWriter (BF16 projections, F32 norms, hand-written index, minimal config.json): every GGUF-named slot is populated after renaming, values round-trip exactly on bf16-representable fixtures, an unmapped INT64 vision-tower tensor is neither delivered nor allowed to trip the engine's fail-fast pre-scan, and Require(BF16) reaches the engine (native bf16 storage for projections, dense FP32 for the norms). Co-Authored-By: Claude Fable 5.1 --- .../GemmaSafeTensorsLoaderFixtureTest.kt | 204 ++++++++++++++++++ 1 file changed, 204 insertions(+) create mode 100644 llm-inference/gemma/src/jvmTest/kotlin/sk/ainet/models/gemma/GemmaSafeTensorsLoaderFixtureTest.kt diff --git a/llm-inference/gemma/src/jvmTest/kotlin/sk/ainet/models/gemma/GemmaSafeTensorsLoaderFixtureTest.kt b/llm-inference/gemma/src/jvmTest/kotlin/sk/ainet/models/gemma/GemmaSafeTensorsLoaderFixtureTest.kt new file mode 100644 index 00000000..661aa70b --- /dev/null +++ b/llm-inference/gemma/src/jvmTest/kotlin/sk/ainet/models/gemma/GemmaSafeTensorsLoaderFixtureTest.kt @@ -0,0 +1,204 @@ +package sk.ainet.models.gemma + +import java.nio.file.Files +import java.nio.file.Path +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertFalse +import kotlin.test.assertNotNull +import kotlin.test.assertTrue +import kotlinx.coroutines.runBlocking +import kotlinx.io.buffered +import kotlinx.io.files.Path as IoPath +import kotlinx.io.files.SystemFileSystem +import sk.ainet.context.DirectCpuExecutionContext +import sk.ainet.io.safetensors.SafeTensorsWriter +import sk.ainet.lang.tensor.data.Bf16DenseTensorData +import sk.ainet.lang.tensor.data.FloatArrayTensorData +import sk.ainet.lang.types.BF16 +import sk.ainet.lang.types.DTypePolicy +import sk.ainet.lang.types.FP32 + +/** + * End-to-end fixture for the engine-backed [GemmaSafeTensorsLoader] (SKaiNET#1246): a synthetic + * 2-shard HF checkpoint in a temp dir — BF16 projections, F32 norms, a hand-written + * `model.safetensors.index.json`, and the minimal `config.json` the loader reads. + * + * Covers what the collapse onto `ShardedSafeTensorsParametersLoader` must preserve: every + * GGUF-named slot is populated after renaming, values round-trip (the fixture values are exactly + * bf16-representable, so equality is exact), an unmapped tensor in the shards is neither delivered + * nor allowed to trip the engine's fail-fast dtype pre-scan (it is INT64 — it would be rejected + * under FP32 if the family allowlist did not exempt it), and the dtype policy reaches the engine + * (`Require(BF16)` keeps native storage). The >2 GiB PLE size guard is not exercised — the + * threshold is a private constant and a fixture that size is not a unit test. + */ +class GemmaSafeTensorsLoaderFixtureTest { + + private val ctx = DirectCpuExecutionContext() + + // Tiny geometry: hidden 8, 2 heads x head_dim 4, 1 kv head, ffn 16, vocab 16, PLE dim 4. + private val hidden = 8 + private val headDim = 4 + private val heads = 2 + private val kvHeads = 1 + private val ffn = 16 + private val vocab = 16 + private val pleDim = 4 + + @Test + fun `every mapped slot is present and values round-trip through both shards`() { + withFixture { indexPath -> + val weights = runBlocking { GemmaSafeTensorsLoader(indexPath).loadToMap(ctx, FP32::class) } + val t = weights.tensors + + val expectedSlots = listOf( + GemmaTensorNames.TOKEN_EMBEDDINGS, GemmaTensorNames.OUTPUT_NORM, GemmaTensorNames.OUTPUT_WEIGHT, + GemmaTensorNames.PER_LAYER_TOKEN_EMBD, GemmaTensorNames.PER_LAYER_MODEL_PROJ, + GemmaTensorNames.PER_LAYER_PROJ_NORM, + GemmaTensorNames.inputLayernorm(0), GemmaTensorNames.attnQ(0), GemmaTensorNames.attnK(0), + GemmaTensorNames.attnV(0), GemmaTensorNames.attnOut(0), GemmaTensorNames.postAttentionNorm(0), + GemmaTensorNames.postAttentionLayernorm(0), GemmaTensorNames.postFfwNorm(0), + GemmaTensorNames.ffnGate(0), GemmaTensorNames.ffnUp(0), GemmaTensorNames.ffnDown(0), + GemmaTensorNames.layerOutputScale(0), GemmaTensorNames.pleInpGate(0), GemmaTensorNames.pleProj(0), + GemmaTensorNames.plePostNorm(0), GemmaTensorNames.attnQNorm(0), GemmaTensorNames.attnKNorm(0), + ) + for (slot in expectedSlots) { + assertNotNull(t[slot], "slot $slot must be populated after HF -> GGUF renaming; got ${t.keys}") + } + assertEquals(expectedSlots.toSet(), t.keys, "no extra or missing slots") + + // Weight tying: the output weight IS the token embedding. + assertTrue(t[GemmaTensorNames.OUTPUT_WEIGHT] === t[GemmaTensorNames.TOKEN_EMBEDDINGS]) + + // Values: shard 1 (BF16 embedding, F32 norm) and shard 2 (BF16 projection, F32 norm). + assertEquals(listOf(vocab, hidden), t[GemmaTensorNames.TOKEN_EMBEDDINGS]!!.shape.dimensions.toList()) + assertEquals(values(vocab * hidden, EMBED_BASE).toList(), floats(t[GemmaTensorNames.TOKEN_EMBEDDINGS]!!).toList()) + assertEquals(values(hidden, NORM_BASE).toList(), floats(t[GemmaTensorNames.OUTPUT_NORM]!!).toList()) + assertEquals(listOf(heads * headDim, hidden), t[GemmaTensorNames.attnQ(0)]!!.shape.dimensions.toList()) + assertEquals(values(heads * headDim * hidden, Q_BASE).toList(), floats(t[GemmaTensorNames.attnQ(0)]!!).toList()) + assertEquals(values(hidden, LN_BASE).toList(), floats(t[GemmaTensorNames.inputLayernorm(0)]!!).toList()) + assertEquals(listOf(1), t[GemmaTensorNames.layerOutputScale(0)]!!.shape.dimensions.toList()) + + // Default policy (Any) widens BF16 to dense FP32. + assertTrue(t[GemmaTensorNames.attnQ(0)]!!.data is FloatArrayTensorData<*>, "Any policy widens BF16") + + // The vision-tower decoy never reached the family (not delivered, not pre-scan-rejected). + assertFalse(t.keys.any { "vision" in it }) + } + } + + @Test + fun `Require BF16 keeps native storage for the projections and leaves F32 norms dense`() { + withFixture { indexPath -> + val weights = runBlocking { + GemmaSafeTensorsLoader(indexPath, dtypePolicy = DTypePolicy.Require(BF16)).loadToMap(ctx, FP32::class) + } + val q = weights.tensors[GemmaTensorNames.attnQ(0)]!! + assertTrue(q.data is Bf16DenseTensorData, "Require(BF16) must keep the projection bf16-native, got ${q.data::class.simpleName}") + // Values still read back exactly (bf16-representable fixture). + assertEquals(values(heads * headDim * hidden, Q_BASE).toList(), floats(q).toList()) + val norm = weights.tensors[GemmaTensorNames.inputLayernorm(0)]!! + assertTrue(norm.data is FloatArrayTensorData<*>, "F32 norms stay dense FP32") + } + } + + // --- fixture --------------------------------------------------------------- + + private fun withFixture(block: (indexPath: String) -> Unit) { + val dir = Files.createTempDirectory("gemma-st-fixture") + try { + val shard1 = "model-00001-of-00002.safetensors" + val shard2 = "model-00002-of-00002.safetensors" + val names1 = mutableListOf() + val names2 = mutableListOf() + + writeShard(dir.resolve(shard1)) { + bf16("model.language_model.embed_tokens.weight", listOf(vocab, hidden), EMBED_BASE, names1) + f32("model.language_model.norm.weight", listOf(hidden), NORM_BASE, names1) + bf16("model.language_model.embed_tokens_per_layer.weight", listOf(vocab, 1 * pleDim), 0.5f, names1) + bf16("model.language_model.per_layer_model_projection.weight", listOf(1 * pleDim, hidden), 0.25f, names1) + f32("model.language_model.per_layer_projection_norm.weight", listOf(pleDim), 1.5f, names1) + } + writeShard(dir.resolve(shard2)) { + val l = "model.language_model.layers.0" + f32("$l.input_layernorm.weight", listOf(hidden), LN_BASE, names2) + bf16("$l.self_attn.q_proj.weight", listOf(heads * headDim, hidden), Q_BASE, names2) + bf16("$l.self_attn.k_proj.weight", listOf(kvHeads * headDim, hidden), 0.125f, names2) + bf16("$l.self_attn.v_proj.weight", listOf(kvHeads * headDim, hidden), 0.25f, names2) + bf16("$l.self_attn.o_proj.weight", listOf(hidden, heads * headDim), 0.375f, names2) + f32("$l.post_attention_layernorm.weight", listOf(hidden), 2.0f, names2) + f32("$l.pre_feedforward_layernorm.weight", listOf(hidden), 2.5f, names2) + f32("$l.post_feedforward_layernorm.weight", listOf(hidden), 3.0f, names2) + bf16("$l.mlp.gate_proj.weight", listOf(ffn, hidden), 0.5f, names2) + bf16("$l.mlp.up_proj.weight", listOf(ffn, hidden), 0.625f, names2) + bf16("$l.mlp.down_proj.weight", listOf(hidden, ffn), 0.75f, names2) + f32("$l.layer_scalar", listOf(1), 4.0f, names2) + bf16("$l.per_layer_input_gate.weight", listOf(pleDim, hidden), 0.875f, names2) + bf16("$l.per_layer_projection.weight", listOf(hidden, pleDim), 1.0f, names2) + f32("$l.post_per_layer_input_norm.weight", listOf(pleDim), 3.5f, names2) + f32("$l.self_attn.q_norm.weight", listOf(headDim), 4.5f, names2) + f32("$l.self_attn.k_norm.weight", listOf(headDim), 5.0f, names2) + // Unmapped decoy: an INT64 tensor the family never asks for. If the allowlist + // did not exempt it, the engine's fail-fast pre-scan would reject the whole load. + tensorI64("model.vision_tower.patch_embedding.position_ids", listOf(4L), LongArray(4) { it.toLong() }) + names2 += "model.vision_tower.patch_embedding.position_ids" + } + + val weightMap = (names1.map { "\"$it\":\"$shard1\"" } + names2.map { "\"$it\":\"$shard2\"" }) + .joinToString(",") + val totalSize = Files.size(dir.resolve(shard1)) + Files.size(dir.resolve(shard2)) + Files.writeString( + dir.resolve("model.safetensors.index.json"), + """{"metadata":{"total_size":$totalSize},"weight_map":{$weightMap}}""" + ) + Files.writeString( + dir.resolve("config.json"), + """ + { + "model_type": "gemma4", + "architectures": ["Gemma4ForConditionalGeneration"], + "text_config": { + "num_hidden_layers": 1, + "hidden_size": $hidden, + "num_attention_heads": $heads, + "num_key_value_heads": $kvHeads, + "head_dim": $headDim, + "vocab_size": $vocab, + "intermediate_size": $ffn, + "hidden_size_per_layer_input": $pleDim + } + } + """.trimIndent() + ) + block(dir.resolve("model.safetensors.index.json").toString()) + } finally { + dir.toFile().deleteRecursively() + } + } + + private fun writeShard(path: Path, block: SafeTensorsWriter.() -> Unit) { + SystemFileSystem.sink(IoPath(path.toString())).buffered().use { sink -> + SafeTensorsWriter.write(sink, block) + } + } + + private fun SafeTensorsWriter.bf16(name: String, shape: List, base: Float, names: MutableList) { + tensorBF16(name, shape.map { it.toLong() }, values(shape.fold(1) { a, b -> a * b }, base)); names += name + } + + private fun SafeTensorsWriter.f32(name: String, shape: List, base: Float, names: MutableList) { + tensorF32(name, shape.map { it.toLong() }, values(shape.fold(1) { a, b -> a * b }, base)); names += name + } + + /** bf16-exact values: base + i/8 stays within 7 mantissa bits for these tiny extents. */ + private fun values(n: Int, base: Float) = FloatArray(n) { base + (it % 16) * 0.125f } + + private fun floats(t: sk.ainet.lang.tensor.Tensor): FloatArray = t.data.copyToFloatArray() + + private companion object { + const val EMBED_BASE = 1.0f + const val NORM_BASE = 2.0f + const val Q_BASE = -1.0f + const val LN_BASE = 0.5f + } +} From d7a270e0ffabcbad915ce5fc0ab0a45677400940 Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Wed, 2 Sep 2026 13:24:44 +0200 Subject: [PATCH 3/3] build: pin the skainet BOM to the published 0.53.0 Brings the engine's billion-parameter export fixes (SKaiNET#1247) and the sharded SafeTensors ParametersLoader (SKaiNET#1246) into the resolved artifact set; the gemma family's loader collapse onto it follows. Co-Authored-By: Claude Fable 5.1 --- gradle/libs.versions.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gradle/libs.versions.toml b/gradle/libs.versions.toml index d524afb0..ff27db84 100644 --- a/gradle/libs.versions.toml +++ b/gradle/libs.versions.toml @@ -1,5 +1,5 @@ [versions] -skainet = "0.52.0" +skainet = "0.53.0" agp = "9.4.0" jacksonDatabind = "2.22.2" jsonSchemaValidator = "3.0.7"