diff --git a/llm-inference/apertus/src/commonMain/kotlin/sk/ainet/models/apertus/ApertusSafeTensorsLoader.kt b/llm-inference/apertus/src/commonMain/kotlin/sk/ainet/models/apertus/ApertusSafeTensorsLoader.kt index 90c02e04..9527cdcd 100644 --- a/llm-inference/apertus/src/commonMain/kotlin/sk/ainet/models/apertus/ApertusSafeTensorsLoader.kt +++ b/llm-inference/apertus/src/commonMain/kotlin/sk/ainet/models/apertus/ApertusSafeTensorsLoader.kt @@ -2,30 +2,39 @@ package sk.ainet.models.apertus 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.StreamingSafeTensorsReader import sk.ainet.io.safetensors.StreamingSafeTensorInfo import sk.ainet.io.safetensors.readTextFile import sk.ainet.io.RandomAccessSource import sk.ainet.lang.tensor.Shape import sk.ainet.lang.tensor.Tensor +import sk.ainet.lang.tensor.data.FloatArrayTensorData import sk.ainet.lang.types.DType +import sk.ainet.lang.types.DTypePolicy import sk.ainet.lang.types.FP32 import kotlin.reflect.KClass /** * Loads Apertus weights from HuggingFace SafeTensors format (sharded). * - * Handles: - * - HuggingFace → GGUF tensor name mapping - * - BF16/F16/F32 dequantization to FP32 - * - Shape normalization ([1, dim] norms → [dim]) - * - Scalar tensor extraction for xIELU parameters - * - Tied word embeddings (output.weight = token_embd.weight) + * Engine delegation (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, the `[1, dim]` norm + * shape normalization, the xIELU scalar extraction, and tied embeddings — 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 ApertusSafeTensorsLoader( - private val indexPath: String + private val indexPath: String, + private val dtypePolicy: DTypePolicy = DTypePolicy.Any, ) { /** @@ -46,30 +55,23 @@ public class ApertusSafeTensorsLoader( val metadata = ApertusConfigParser.parse(configJson) - val reader = StreamingShardedSafeTensorsReader.openFromIndex(indexPath) - - return reader.use { - loadFromReader(ctx, dtype, it, metadata) - } - } + // Family policy as the engine loader's filter: only the tensors this family maps + // (anything unmapped never materializes and is exempt from the fail-fast dtype pre-scan). + val wanted = wantedHfNames(metadata.blockCount) + val loader = ShardedSafeTensorsParametersLoader.withPolicy( + indexPath = indexPath, + policy = dtypePolicy, + tensorFilter = { info: ShardedTensorInfo -> info.name in wanted }, + ) + val tensorsByHfName = linkedMapOf>() + loader.load(ctx, dtype) { name, tensor -> tensorsByHfName[name] = tensor } - private fun loadFromReader( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - metadata: ApertusModelMetadata - ): ApertusWeights { val tensorsByGgufName = linkedMapOf>() val xieluParams = mutableMapOf() - val tensorsByHfName = reader.tensors.associateBy { it.name } - - // Load global tensors - loadGlobalTensors(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, metadata) - - // Load layer tensors + loadGlobalTensors(ctx, dtype, tensorsByHfName, tensorsByGgufName, metadata) for (layer in 0 until metadata.blockCount) { - loadLayerTensors(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, xieluParams, layer) + loadLayerTensors(ctx, dtype, tensorsByHfName, tensorsByGgufName, xieluParams, layer) } return ApertusWeights( @@ -82,31 +84,24 @@ public class ApertusSafeTensorsLoader( private fun loadGlobalTensors( ctx: ExecutionContext, dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, + tensorsByHfName: Map>, tensorsByGgufName: MutableMap>, metadata: ApertusModelMetadata ) { // Token embeddings - val embedTokens = tensorsByHfName[HF_EMBED_TOKENS] - ?: error("Missing tensor: $HF_EMBED_TOKENS") - val embedTensor = loadAndConvertTensor(ctx, dtype, reader, embedTokens, transpose = false) + val embedTensor = normalizeNormShape(ctx, dtype, required(tensorsByHfName, HF_EMBED_TOKENS)) tensorsByGgufName[ApertusTensorNames.TOKEN_EMBEDDINGS] = embedTensor // Output norm - val norm = tensorsByHfName[HF_OUTPUT_NORM] - ?: error("Missing tensor: $HF_OUTPUT_NORM") - val normTensor = loadAndConvertTensor(ctx, dtype, reader, norm, transpose = false) - tensorsByGgufName[ApertusTensorNames.OUTPUT_NORM] = normTensor + tensorsByGgufName[ApertusTensorNames.OUTPUT_NORM] = + normalizeNormShape(ctx, dtype, required(tensorsByHfName, HF_OUTPUT_NORM)) // Output weight (may be tied to embeddings) val lmHead = tensorsByHfName[HF_LM_HEAD] if (lmHead != null) { - tensorsByGgufName[ApertusTensorNames.OUTPUT_WEIGHT] = - loadAndConvertTensor(ctx, dtype, reader, lmHead, transpose = false) + tensorsByGgufName[ApertusTensorNames.OUTPUT_WEIGHT] = normalizeNormShape(ctx, dtype, lmHead) } else if (metadata.tiedEmbeddings) { tensorsByGgufName[ApertusTensorNames.OUTPUT_WEIGHT] = embedTensor - println(" Tied: ${ApertusTensorNames.OUTPUT_WEIGHT} → ${ApertusTensorNames.TOKEN_EMBEDDINGS}") } else { error("Missing lm_head.weight and tie_word_embeddings is false") } @@ -115,146 +110,54 @@ public class ApertusSafeTensorsLoader( private fun loadLayerTensors( ctx: ExecutionContext, dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, + tensorsByHfName: Map>, tensorsByGgufName: MutableMap>, xieluParams: MutableMap, layer: Int ) { - // Attention norm (Apertus uses "attention_layernorm") - loadRequiredTensor(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnNorm(layer), ApertusTensorNames.attnNorm(layer), transpose = false) - - // QKV + output projections - loadRequiredTensor(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnQ(layer), ApertusTensorNames.attnQ(layer)) - loadRequiredTensor(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnK(layer), ApertusTensorNames.attnK(layer)) - loadRequiredTensor(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnV(layer), ApertusTensorNames.attnV(layer)) - loadRequiredTensor(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnO(layer), ApertusTensorNames.attnOut(layer)) - - // QK-norm weights - loadRequiredTensor(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfQNorm(layer), ApertusTensorNames.attnQNorm(layer), transpose = false) - loadRequiredTensor(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfKNorm(layer), ApertusTensorNames.attnKNorm(layer), transpose = false) - - // FFN norm (Apertus uses "feedforward_layernorm") - loadRequiredTensor(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfFfnNorm(layer), ApertusTensorNames.ffnNorm(layer), transpose = false) - - // MLP weights (ungated: up + down only, no gate_proj) - loadRequiredTensor(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfMlpUp(layer), ApertusTensorNames.ffnUp(layer)) - loadRequiredTensor(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfMlpDown(layer), ApertusTensorNames.ffnDown(layer)) - - // xIELU scalar parameters - val alphaP = loadScalarParam(reader, tensorsByHfName, hfXieluAlphaP(layer)) - val alphaN = loadScalarParam(reader, tensorsByHfName, hfXieluAlphaN(layer)) - val beta = loadScalarParam(reader, tensorsByHfName, hfXieluBeta(layer)) - val eps = loadScalarParam(reader, tensorsByHfName, hfXieluEps(layer)) + for ((hfName, ggufName) in layerSlots(layer)) { + tensorsByGgufName[ggufName] = normalizeNormShape(ctx, dtype, required(tensorsByHfName, hfName)) + } + + // xIELU scalar parameters (shape [] or [1]; consumed as plain floats). xieluParams[layer] = ApertusXIELUParams( - alphaP = alphaP, - alphaN = alphaN, - beta = beta, - eps = eps + alphaP = scalar(tensorsByHfName, hfXieluAlphaP(layer)), + alphaN = scalar(tensorsByHfName, hfXieluAlphaN(layer)), + beta = scalar(tensorsByHfName, hfXieluBeta(layer)), + eps = scalar(tensorsByHfName, hfXieluEps(layer)) ) } - private fun loadRequiredTensor( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, - tensorsByGgufName: MutableMap>, - hfName: String, - ggufName: String, - transpose: Boolean = false - ) { - val info = tensorsByHfName[hfName] - ?: error("Missing required tensor: $hfName") - val tensor = loadAndConvertTensor(ctx, dtype, reader, info, transpose) - tensorsByGgufName[ggufName] = tensor - println(" Loaded: $hfName (${info.dtype} ${info.shape}) → $ggufName") + private fun required(tensorsByHfName: Map>, hfName: String): Tensor = + tensorsByHfName[hfName] ?: error("Missing required tensor: $hfName") + + private fun scalar(tensorsByHfName: Map>, hfName: String): Float { + val tensor = tensorsByHfName[hfName] ?: error("Missing scalar tensor: $hfName") + return tensor.data.copyToFloatArray()[0] } /** - * Load a scalar parameter from a tensor with shape [] or [1]. + * Some exports store norms as `[1, dim]`; the runtime expects `[dim]`. A widened tensor is + * re-wrapped over the same buffer (no copy); a keep-native tensor is widened for this one + * defensive edge case, since the norm is tiny. */ - private fun loadScalarParam( - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, - hfName: String - ): Float { - val info = tensorsByHfName[hfName] - ?: error("Missing scalar tensor: $hfName") - val bytes = reader.loadTensorData(info) - return when (info.dtype.uppercase()) { - "BF16" -> dequantBF16FromBytes(bytes)[0] - "F16" -> dequantF16FromBytes(bytes)[0] - "F32" -> bytesToFloatArray(bytes)[0] - else -> error("Unsupported dtype for scalar: ${info.dtype}") - } - } - @Suppress("UNCHECKED_CAST") - private fun loadAndConvertTensor( + private fun normalizeNormShape( ctx: ExecutionContext, dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - info: ShardedTensorInfo, - transpose: Boolean = false + tensor: Tensor ): Tensor { - val bytes = reader.loadTensorData(info) - val shape = normalizeNormShape(info.shape) - - val floats = when (info.dtype.uppercase()) { - "BF16" -> dequantBF16FromBytes(bytes) - "F16" -> dequantF16FromBytes(bytes) - "F32" -> bytesToFloatArray(bytes) - else -> error("Unsupported SafeTensors dtype: ${info.dtype}") - } - - 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) + val dims = tensor.shape.dimensions + if (dims.size != 2 || dims[0] != 1) return tensor + val newShape = Shape(dims[1]) + val data = tensor.data + return if (data is FloatArrayTensorData<*>) { + ctx.wrapFloatArray(newShape, dtype, (data as FloatArrayTensorData).buffer) } else { - ctx.fromFloatArray(shape, dtype, floats) + ctx.fromFloatArray(newShape, dtype, data.copyToFloatArray()) } } - private fun normalizeNormShape(shape: List): Shape { - return if (shape.size == 2 && shape[0] == 1L) { - Shape(shape[1].toInt()) - } else if (shape.isEmpty()) { - Shape(1) // scalar → [1] - } else { - Shape(*shape.map { it.toInt() }.toIntArray()) - } - } - - // ========== Byte Conversion Helpers ========== - - private fun bytesToFloatArray(bytes: ByteArray): FloatArray = DequantOps.bytesToFloatArray(bytes) - private fun dequantF16FromBytes(bytes: ByteArray): FloatArray = DequantOps.dequantF16FromBytes(bytes) - private fun dequantBF16FromBytes(bytes: ByteArray): FloatArray = DequantOps.dequantBF16FromBytes(bytes) - - 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] - } - } - return out - } - // ========== HuggingFace Tensor Name Constants ========== private companion object { @@ -262,6 +165,34 @@ public class ApertusSafeTensorsLoader( const val HF_OUTPUT_NORM = "model.norm.weight" const val HF_LM_HEAD = "lm_head.weight" + /** 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); add(HF_LM_HEAD) + for (layer in 0 until blockCount) { + addAll(layerSlots(layer).map { it.first }) + add(hfXieluAlphaP(layer)); add(hfXieluAlphaN(layer)); add(hfXieluBeta(layer)); add(hfXieluEps(layer)) + } + } + + /** Per-layer HF → GGUF slot pairs; every one is required. */ + fun layerSlots(layer: Int): List> = listOf( + // Attention norm (Apertus uses "attention_layernorm") + hfAttnNorm(layer) to ApertusTensorNames.attnNorm(layer), + // QKV + output projections + hfAttnQ(layer) to ApertusTensorNames.attnQ(layer), + hfAttnK(layer) to ApertusTensorNames.attnK(layer), + hfAttnV(layer) to ApertusTensorNames.attnV(layer), + hfAttnO(layer) to ApertusTensorNames.attnOut(layer), + // QK-norm weights + hfQNorm(layer) to ApertusTensorNames.attnQNorm(layer), + hfKNorm(layer) to ApertusTensorNames.attnKNorm(layer), + // FFN norm (Apertus uses "feedforward_layernorm") + hfFfnNorm(layer) to ApertusTensorNames.ffnNorm(layer), + // MLP weights (ungated: up + down only, no gate_proj) + hfMlpUp(layer) to ApertusTensorNames.ffnUp(layer), + hfMlpDown(layer) to ApertusTensorNames.ffnDown(layer), + ) + fun hfAttnNorm(layer: Int) = "model.layers.$layer.attention_layernorm.weight" fun hfAttnQ(layer: Int) = "model.layers.$layer.self_attn.q_proj.weight" fun hfAttnK(layer: Int) = "model.layers.$layer.self_attn.k_proj.weight" diff --git a/llm-inference/apertus/src/jvmTest/kotlin/sk/ainet/models/apertus/ApertusSafeTensorsLoaderFixtureTest.kt b/llm-inference/apertus/src/jvmTest/kotlin/sk/ainet/models/apertus/ApertusSafeTensorsLoaderFixtureTest.kt new file mode 100644 index 00000000..6517001c --- /dev/null +++ b/llm-inference/apertus/src/jvmTest/kotlin/sk/ainet/models/apertus/ApertusSafeTensorsLoaderFixtureTest.kt @@ -0,0 +1,196 @@ +package sk.ainet.models.apertus + +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 [ApertusSafeTensorsLoader] (SKaiNET#1246): a synthetic + * 2-shard HF checkpoint in a temp dir — BF16 projections, F32 norms (one stored `[1, dim]` to + * exercise the family-side shape normalization), scalar xIELU parameters, 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 exactly (bf16-representable + * fixture), tied embeddings, the xIELU scalars arrive as floats, an unmapped INT64 decoy is + * neither delivered nor allowed to trip the engine's fail-fast dtype pre-scan, and the dtype + * policy reaches the engine (`Require(BF16)` keeps native storage). + */ +class ApertusSafeTensorsLoaderFixtureTest { + + private val ctx = DirectCpuExecutionContext() + + private val hidden = 8 + private val headDim = 4 + private val heads = 2 + private val kvHeads = 1 + private val ffn = 16 + private val vocab = 16 + + @Test + fun `every mapped slot is present, norms are normalized and values round-trip`() { + withFixture { indexPath -> + val weights = runBlocking { ApertusSafeTensorsLoader(indexPath).loadToMap(ctx, FP32::class) } + val t = weights.tensors + + val expectedSlots = listOf( + ApertusTensorNames.TOKEN_EMBEDDINGS, ApertusTensorNames.OUTPUT_NORM, ApertusTensorNames.OUTPUT_WEIGHT, + ApertusTensorNames.attnNorm(0), ApertusTensorNames.attnQ(0), ApertusTensorNames.attnK(0), + ApertusTensorNames.attnV(0), ApertusTensorNames.attnOut(0), ApertusTensorNames.attnQNorm(0), + ApertusTensorNames.attnKNorm(0), ApertusTensorNames.ffnNorm(0), ApertusTensorNames.ffnUp(0), + ApertusTensorNames.ffnDown(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") + + // Tied embeddings (no lm_head.weight in the shards, tie_word_embeddings=true). + assertTrue(t[ApertusTensorNames.OUTPUT_WEIGHT] === t[ApertusTensorNames.TOKEN_EMBEDDINGS]) + + // Values and shapes across both shards. + assertEquals(listOf(vocab, hidden), t[ApertusTensorNames.TOKEN_EMBEDDINGS]!!.shape.dimensions.toList()) + assertEquals(values(vocab * hidden, EMBED_BASE).toList(), floats(t[ApertusTensorNames.TOKEN_EMBEDDINGS]!!).toList()) + assertEquals(values(heads * headDim * hidden, Q_BASE).toList(), floats(t[ApertusTensorNames.attnQ(0)]!!).toList()) + + // The attention norm was stored [1, hidden]; the runtime slot must be [hidden] with the same values. + assertEquals(listOf(hidden), t[ApertusTensorNames.attnNorm(0)]!!.shape.dimensions.toList()) + assertEquals(values(hidden, LN_BASE).toList(), floats(t[ApertusTensorNames.attnNorm(0)]!!).toList()) + + // xIELU scalars arrive as plain floats (shape [] and [1] both accepted). + val xielu = weights.xieluParams[0]!! + assertEquals(0.8f, xielu.alphaP) + assertEquals(0.5f, xielu.alphaN) + assertEquals(0.75f, xielu.beta) + assertEquals(0.125f, xielu.eps) + + // Default policy (Any) widens BF16 to dense FP32. + assertTrue(t[ApertusTensorNames.attnQ(0)]!!.data is FloatArrayTensorData<*>, "Any policy widens BF16") + + // The decoy never reached the family. + assertFalse(t.keys.any { "rotary" in it }) + } + } + + @Test + fun `Require BF16 keeps native storage for the projections and leaves F32 norms dense`() { + withFixture { indexPath -> + val weights = runBlocking { + ApertusSafeTensorsLoader(indexPath, dtypePolicy = DTypePolicy.Require(BF16)).loadToMap(ctx, FP32::class) + } + val q = weights.tensors[ApertusTensorNames.attnQ(0)]!! + assertTrue(q.data is Bf16DenseTensorData, "Require(BF16) must keep the projection bf16-native, got ${q.data::class.simpleName}") + assertEquals(values(heads * headDim * hidden, Q_BASE).toList(), floats(q).toList()) + val norm = weights.tensors[ApertusTensorNames.ffnNorm(0)]!! + assertTrue(norm.data is FloatArrayTensorData<*>, "F32 norms stay dense FP32") + } + } + + // --- fixture --------------------------------------------------------------- + + private fun withFixture(block: (indexPath: String) -> Unit) { + val dir = Files.createTempDirectory("apertus-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.embed_tokens.weight", listOf(vocab, hidden), EMBED_BASE, names1) + f32("model.norm.weight", listOf(hidden), NORM_BASE, names1) + } + writeShard(dir.resolve(shard2)) { + val l = "model.layers.0" + // Stored [1, hidden] on purpose: the family normalizes it to [hidden]. + f32("$l.attention_layernorm.weight", listOf(1, 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.self_attn.q_norm.weight", listOf(headDim), 4.5f, names2) + f32("$l.self_attn.k_norm.weight", listOf(headDim), 5.0f, names2) + f32("$l.feedforward_layernorm.weight", listOf(hidden), 2.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) + // xIELU scalars: two stored as rank-0, two as [1]. + tensorF32("$l.mlp.act_fn.alpha_p", emptyList(), floatArrayOf(0.8f)); names2 += "$l.mlp.act_fn.alpha_p" + tensorF32("$l.mlp.act_fn.alpha_n", emptyList(), floatArrayOf(0.5f)); names2 += "$l.mlp.act_fn.alpha_n" + tensorBF16("$l.mlp.act_fn.beta", listOf(1L), floatArrayOf(0.75f)); names2 += "$l.mlp.act_fn.beta" + tensorF32("$l.mlp.act_fn.eps", listOf(1L), floatArrayOf(0.125f)); names2 += "$l.mlp.act_fn.eps" + // 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("$l.self_attn.rotary_emb.inv_freq_ids", listOf(4L), LongArray(4) { it.toLong() }) + names2 += "$l.self_attn.rotary_emb.inv_freq_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": "apertus", + "hidden_size": $hidden, + "num_hidden_layers": 1, + "num_attention_heads": $heads, + "num_key_value_heads": $kvHeads, + "head_dim": $headDim, + "intermediate_size": $ffn, + "vocab_size": $vocab, + "tie_word_embeddings": true + } + """.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 + } +} diff --git a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nSafeTensorsWeightLoader.kt b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nSafeTensorsWeightLoader.kt index a9e821b4..436deada 100644 --- a/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nSafeTensorsWeightLoader.kt +++ b/llm-inference/gemma3n/src/commonMain/kotlin/sk/ainet/models/gemma3n/Gemma3nSafeTensorsWeightLoader.kt @@ -1,13 +1,13 @@ package sk.ainet.models.gemma3n 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 @@ -20,19 +20,31 @@ import kotlin.reflect.KClass * Supports sharded models (multiple .safetensors files with index.json). * * Key differences from GGUF: - * - Uses BF16 dtype (needs conversion to FP32) * - Different tensor name format (model.language_model.layers.X.* vs blk.X.*) * - Weight tying: embed_tokens.weight is reused for output projection + * + * Engine delegation (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-table 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 Gemma3nSafeTensorsWeightLoader( - private val indexPath: String + private val indexPath: String, + private val dtypePolicy: DTypePolicy = DTypePolicy.Any, ) { /** * Load weights into a map, mapping HuggingFace names to GGUF-style names. * * @param ctx Execution context for tensor operations - * @param dtype Target dtype (FP32 or FP16) + * @param dtype Target dtype (FP32) * @return Gemma3nWeights with mapped tensor names */ public suspend fun loadToMap( @@ -51,31 +63,28 @@ public class Gemma3nSafeTensorsWeightLoader( val metadata = Gemma3nConfigParser.parseFromJson(configJson) - // Open sharded reader - val reader = StreamingShardedSafeTensorsReader.openFromIndex(indexPath) - - return reader.use { - loadFromReader(ctx, dtype, it, metadata) - } - } + // Family policy 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 path — the + // engine returns each tensor as one array, which the JVM caps at Int.MAX_VALUE bytes; + // an oversized table is left absent, which disables PLE downstream instead of failing. + val wanted = wantedHfNames(metadata.blockCount) + val loader = ShardedSafeTensorsParametersLoader.withPolicy( + indexPath = indexPath, + policy = dtypePolicy, + tensorFilter = { info: ShardedTensorInfo -> + info.name in wanted && (info.name != HF_PER_LAYER_TOKEN_EMBD || 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: Gemma3nModelMetadata - ): Gemma3nWeights { val tensorsByGgufName = linkedMapOf>() - - // Build lookup by HuggingFace name - val tensorsByHfName = reader.tensors.associateBy { it.name } - - // Load global tensors - loadGlobalTensors(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, metadata) - - // Load layer tensors + loadGlobalTensors(tensorsByHfName, tensorsByGgufName) for (layer in 0 until metadata.blockCount) { - loadLayerTensors(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, metadata, layer) + for ((hfName, ggufName) in layerSlots(layer)) { + tensorsByHfName[hfName]?.let { tensorsByGgufName[ggufName] = it } + } } return Gemma3nWeights( @@ -85,222 +94,34 @@ public class Gemma3nSafeTensorsWeightLoader( } private fun loadGlobalTensors( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, - tensorsByGgufName: MutableMap>, - metadata: Gemma3nModelMetadata + tensorsByHfName: Map>, + tensorsByGgufName: MutableMap> ) { - // Token embeddings - do NOT transpose, keep as [vocab_size, embedding_dim] - val embedTokens = tensorsByHfName[HF_EMBED_TOKENS] + // Token embeddings, kept as [vocab_size, embedding_dim] + val embedTensor = tensorsByHfName[HF_EMBED_TOKENS] ?: error("Missing tensor: $HF_EMBED_TOKENS") - val embedTensor = loadAndConvertTensor(ctx, dtype, reader, embedTokens, transpose = false) tensorsByGgufName[Gemma3nTensorNames.TOKEN_EMBEDDINGS] = embedTensor // Output norm - val norm = tensorsByHfName[HF_OUTPUT_NORM] + tensorsByGgufName[Gemma3nTensorNames.OUTPUT_NORM] = tensorsByHfName[HF_OUTPUT_NORM] ?: error("Missing tensor: $HF_OUTPUT_NORM") - val normTensor = loadAndConvertTensor(ctx, dtype, reader, norm, transpose = false) - tensorsByGgufName[Gemma3nTensorNames.OUTPUT_NORM] = normTensor // Output weight (weight tying - reuse embed_tokens) - // Gemma 3n ties the output projection to the embedding weights tensorsByGgufName[Gemma3nTensorNames.OUTPUT_WEIGHT] = embedTensor - // Global AltUp tensors (optional, E4B) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - HF_ALTUP_PROJ, Gemma3nTensorNames.ALTUP_PROJ, - transpose = false - ) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - HF_ALTUP_UNEMBD_PROJ, Gemma3nTensorNames.ALTUP_UNEMBD_PROJ, - transpose = false - ) - - // Global per-layer embedding tensors (E4B) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - HF_PER_LAYER_TOKEN_EMBD, Gemma3nTensorNames.PER_LAYER_TOKEN_EMBD, - transpose = false - ) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - HF_PER_LAYER_MODEL_PROJ, Gemma3nTensorNames.PER_LAYER_MODEL_PROJ, - transpose = false - ) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - HF_PER_LAYER_PROJ_NORM, Gemma3nTensorNames.PER_LAYER_PROJ_NORM, - transpose = false - ) - } - - private fun loadLayerTensors( - ctx: ExecutionContext, - dtype: KClass, - reader: StreamingShardedSafeTensorsReader, - tensorsByHfName: Map, - tensorsByGgufName: MutableMap>, - metadata: Gemma3nModelMetadata, - layer: Int - ) { - // Input layernorm - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfInputLayernorm(layer), Gemma3nTensorNames.inputLayernorm(layer), - transpose = false - ) - - // Attention weights - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnQ(layer), Gemma3nTensorNames.attnQ(layer) - ) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnK(layer), Gemma3nTensorNames.attnK(layer) - ) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnV(layer), Gemma3nTensorNames.attnV(layer) - ) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnO(layer), Gemma3nTensorNames.attnOut(layer) - ) - - // Post-attention layernorm - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPostAttnLayernorm(layer), Gemma3nTensorNames.postAttentionLayernorm(layer), - transpose = false - ) - - // MLP weights - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfMlpGate(layer), Gemma3nTensorNames.ffnGate(layer) - ) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfMlpUp(layer), Gemma3nTensorNames.ffnUp(layer) - ) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfMlpDown(layer), Gemma3nTensorNames.ffnDown(layer) - ) - - // Per-layer projection (optional) - loadTensorIfExists( - ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPerLayerProjection(layer), Gemma3nTensorNames.perLayerInput(layer) - ) - - // E4B per-layer AltUp tensors - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAltupPredictCoef(layer), Gemma3nTensorNames.altupPredictCoef(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAltupCorrectCoef(layer), Gemma3nTensorNames.altupCorrectCoef(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAltupCorrectScale(layer), Gemma3nTensorNames.altupCorrectScale(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAltupRouter(layer), Gemma3nTensorNames.altupRouter(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAltupRouterNorm(layer), Gemma3nTensorNames.altupRouterNorm(layer), transpose = false) - - // E4B additional norms and weights - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnQNorm(layer), Gemma3nTensorNames.attnQNorm(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfAttnKNorm(layer), Gemma3nTensorNames.attnKNorm(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPostAttentionNorm(layer), Gemma3nTensorNames.postAttentionNorm(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPostFfwNorm(layer), Gemma3nTensorNames.postFfwNorm(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfPostNorm(layer), Gemma3nTensorNames.postNorm(layer), transpose = false) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfInputGate(layer), Gemma3nTensorNames.inputGate(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfProj(layer), Gemma3nTensorNames.proj(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfLaurelL(layer), Gemma3nTensorNames.laurelL(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfLaurelR(layer), Gemma3nTensorNames.laurelR(layer)) - loadTensorIfExists(ctx, dtype, reader, tensorsByHfName, tensorsByGgufName, - hfLaurelPostNorm(layer), Gemma3nTensorNames.laurelPostNorm(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 + // Optional globals: AltUp (E4B) and per-layer embedding tensors. + for ((hfName, ggufName) in GLOBAL_OPTIONAL_SLOTS) { + tensorsByHfName[hfName]?.let { tensorsByGgufName[ggufName] = it } } } - @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()) - - // Convert bytes to float array based on dtype - val floats = when (info.dtype.uppercase()) { - "BF16" -> dequantBF16FromBytes(bytes) - "F16" -> dequantF16FromBytes(bytes) - "F32" -> bytesToFloatArray(bytes) - else -> error("Unsupported SafeTensors dtype: ${info.dtype}") - } - - // Transpose 2D tensors from row-major (PyTorch) to our expected format - 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 - } - } - - // ========== Byte Conversion Helpers (delegating to DequantOps) ========== - - private fun bytesToFloatArray(bytes: ByteArray): FloatArray = DequantOps.bytesToFloatArray(bytes) - private fun dequantF16FromBytes(bytes: ByteArray): FloatArray = DequantOps.dequantF16FromBytes(bytes) - private fun dequantBF16FromBytes(bytes: ByteArray): FloatArray = DequantOps.dequantBF16FromBytes(bytes) - - 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] - } - } - return out - } - // ========== HuggingFace Tensor Name Constants ========== private companion object { + // The engine's eager loader path returns each tensor as a single + // array, which the JVM caps at Int.MAX_VALUE bytes. Keep a little headroom. + const val MAX_BYTES_PER_TENSOR: Long = Int.MAX_VALUE.toLong() - 1024L + // Global tensors const val HF_EMBED_TOKENS = "model.language_model.embed_tokens.weight" const val HF_OUTPUT_NORM = "model.language_model.norm.weight" @@ -314,6 +135,55 @@ public class Gemma3nSafeTensorsWeightLoader( const val HF_PER_LAYER_MODEL_PROJ = "model.language_model.per_layer_model_proj.weight" const val HF_PER_LAYER_PROJ_NORM = "model.language_model.per_layer_proj_norm.weight" + /** Optional global HF → GGUF slot pairs. */ + val GLOBAL_OPTIONAL_SLOTS: List> by lazy { + listOf( + HF_ALTUP_PROJ to Gemma3nTensorNames.ALTUP_PROJ, + HF_ALTUP_UNEMBD_PROJ to Gemma3nTensorNames.ALTUP_UNEMBD_PROJ, + HF_PER_LAYER_TOKEN_EMBD to Gemma3nTensorNames.PER_LAYER_TOKEN_EMBD, + HF_PER_LAYER_MODEL_PROJ to Gemma3nTensorNames.PER_LAYER_MODEL_PROJ, + HF_PER_LAYER_PROJ_NORM to Gemma3nTensorNames.PER_LAYER_PROJ_NORM, + ) + } + + /** 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(GLOBAL_OPTIONAL_SLOTS.map { it.first }) + for (layer in 0 until blockCount) addAll(layerSlots(layer).map { it.first }) + } + + /** Per-layer HF → GGUF slot pairs; all optional (E2B vs E4B differ), mapped when present. */ + fun layerSlots(layer: Int): List> = listOf( + hfInputLayernorm(layer) to Gemma3nTensorNames.inputLayernorm(layer), + hfAttnQ(layer) to Gemma3nTensorNames.attnQ(layer), + hfAttnK(layer) to Gemma3nTensorNames.attnK(layer), + hfAttnV(layer) to Gemma3nTensorNames.attnV(layer), + hfAttnO(layer) to Gemma3nTensorNames.attnOut(layer), + hfPostAttnLayernorm(layer) to Gemma3nTensorNames.postAttentionLayernorm(layer), + hfMlpGate(layer) to Gemma3nTensorNames.ffnGate(layer), + hfMlpUp(layer) to Gemma3nTensorNames.ffnUp(layer), + hfMlpDown(layer) to Gemma3nTensorNames.ffnDown(layer), + hfPerLayerProjection(layer) to Gemma3nTensorNames.perLayerInput(layer), + // E4B per-layer AltUp tensors + hfAltupPredictCoef(layer) to Gemma3nTensorNames.altupPredictCoef(layer), + hfAltupCorrectCoef(layer) to Gemma3nTensorNames.altupCorrectCoef(layer), + hfAltupCorrectScale(layer) to Gemma3nTensorNames.altupCorrectScale(layer), + hfAltupRouter(layer) to Gemma3nTensorNames.altupRouter(layer), + hfAltupRouterNorm(layer) to Gemma3nTensorNames.altupRouterNorm(layer), + // E4B additional norms and weights + hfAttnQNorm(layer) to Gemma3nTensorNames.attnQNorm(layer), + hfAttnKNorm(layer) to Gemma3nTensorNames.attnKNorm(layer), + hfPostAttentionNorm(layer) to Gemma3nTensorNames.postAttentionNorm(layer), + hfPostFfwNorm(layer) to Gemma3nTensorNames.postFfwNorm(layer), + hfPostNorm(layer) to Gemma3nTensorNames.postNorm(layer), + hfInputGate(layer) to Gemma3nTensorNames.inputGate(layer), + hfProj(layer) to Gemma3nTensorNames.proj(layer), + hfLaurelL(layer) to Gemma3nTensorNames.laurelL(layer), + hfLaurelR(layer) to Gemma3nTensorNames.laurelR(layer), + hfLaurelPostNorm(layer) to Gemma3nTensorNames.laurelPostNorm(layer), + ) + // Layer tensor name builders fun hfInputLayernorm(layer: Int) = "model.language_model.layers.$layer.input_layernorm.weight" fun hfAttnQ(layer: Int) = "model.language_model.layers.$layer.self_attn.q_proj.weight" diff --git a/llm-inference/gemma3n/src/jvmTest/kotlin/sk/ainet/models/gemma3n/Gemma3nSafeTensorsWeightLoaderFixtureTest.kt b/llm-inference/gemma3n/src/jvmTest/kotlin/sk/ainet/models/gemma3n/Gemma3nSafeTensorsWeightLoaderFixtureTest.kt new file mode 100644 index 00000000..d7741a02 --- /dev/null +++ b/llm-inference/gemma3n/src/jvmTest/kotlin/sk/ainet/models/gemma3n/Gemma3nSafeTensorsWeightLoaderFixtureTest.kt @@ -0,0 +1,204 @@ +package sk.ainet.models.gemma3n + +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.assertNull +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 [Gemma3nSafeTensorsWeightLoader] (SKaiNET#1246): a + * synthetic 2-shard HF checkpoint in a temp dir — BF16 projections, F32 norms, an E2B-style + * layer (no AltUp, no Laurel) plus the PLE globals, a hand-written + * `model.safetensors.index.json`, and the minimal `config.json` the loader reads. + * + * Covers what the collapse onto `ShardedSafeTensorsParametersLoader` must preserve: every + * present GGUF-named slot is populated after renaming and absent optional (E4B) slots stay + * absent, values round-trip exactly (bf16-representable fixture), tied output/embedding, an + * unmapped INT64 decoy is neither delivered nor allowed to trip the engine's fail-fast dtype + * pre-scan, and the dtype policy reaches the engine (`Require(BF16)` keeps native storage). + * The >2 GiB PLE size guard is not exercised — a fixture that size is not a unit test. + */ +class Gemma3nSafeTensorsWeightLoaderFixtureTest { + + private val ctx = DirectCpuExecutionContext() + + 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 `present slots are mapped, absent E4B slots stay absent, values round-trip`() { + withFixture { indexPath -> + val weights = runBlocking { Gemma3nSafeTensorsWeightLoader(indexPath).loadToMap(ctx, FP32::class) } + val t = weights.tensors + + val expectedSlots = listOf( + Gemma3nTensorNames.TOKEN_EMBEDDINGS, Gemma3nTensorNames.OUTPUT_NORM, Gemma3nTensorNames.OUTPUT_WEIGHT, + Gemma3nTensorNames.PER_LAYER_TOKEN_EMBD, Gemma3nTensorNames.PER_LAYER_MODEL_PROJ, + Gemma3nTensorNames.PER_LAYER_PROJ_NORM, + Gemma3nTensorNames.inputLayernorm(0), Gemma3nTensorNames.attnQ(0), Gemma3nTensorNames.attnK(0), + Gemma3nTensorNames.attnV(0), Gemma3nTensorNames.attnOut(0), Gemma3nTensorNames.postAttentionLayernorm(0), + Gemma3nTensorNames.ffnGate(0), Gemma3nTensorNames.ffnUp(0), Gemma3nTensorNames.ffnDown(0), + Gemma3nTensorNames.perLayerInput(0), Gemma3nTensorNames.attnQNorm(0), Gemma3nTensorNames.attnKNorm(0), + Gemma3nTensorNames.postAttentionNorm(0), Gemma3nTensorNames.postFfwNorm(0), Gemma3nTensorNames.postNorm(0), + Gemma3nTensorNames.inputGate(0), Gemma3nTensorNames.proj(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") + + // Optional E4B tensors were not in the shards and must not be synthesized. + assertNull(t[Gemma3nTensorNames.altupRouter(0)]) + assertNull(t[Gemma3nTensorNames.laurelL(0)]) + assertNull(t[Gemma3nTensorNames.ALTUP_PROJ]) + + // Weight tying: the output weight IS the token embedding. + assertTrue(t[Gemma3nTensorNames.OUTPUT_WEIGHT] === t[Gemma3nTensorNames.TOKEN_EMBEDDINGS]) + + assertEquals(listOf(vocab, hidden), t[Gemma3nTensorNames.TOKEN_EMBEDDINGS]!!.shape.dimensions.toList()) + assertEquals(values(vocab * hidden, EMBED_BASE).toList(), floats(t[Gemma3nTensorNames.TOKEN_EMBEDDINGS]!!).toList()) + assertEquals(values(hidden, NORM_BASE).toList(), floats(t[Gemma3nTensorNames.OUTPUT_NORM]!!).toList()) + assertEquals(listOf(heads * headDim, hidden), t[Gemma3nTensorNames.attnQ(0)]!!.shape.dimensions.toList()) + assertEquals(values(heads * headDim * hidden, Q_BASE).toList(), floats(t[Gemma3nTensorNames.attnQ(0)]!!).toList()) + assertEquals(listOf(vocab, pleDim), t[Gemma3nTensorNames.PER_LAYER_TOKEN_EMBD]!!.shape.dimensions.toList()) + + assertTrue(t[Gemma3nTensorNames.attnQ(0)]!!.data is FloatArrayTensorData<*>, "Any policy widens BF16") + 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 { + Gemma3nSafeTensorsWeightLoader(indexPath, dtypePolicy = DTypePolicy.Require(BF16)).loadToMap(ctx, FP32::class) + } + val q = weights.tensors[Gemma3nTensorNames.attnQ(0)]!! + assertTrue(q.data is Bf16DenseTensorData, "Require(BF16) must keep the projection bf16-native, got ${q.data::class.simpleName}") + assertEquals(values(heads * headDim * hidden, Q_BASE).toList(), floats(q).toList()) + val norm = weights.tensors[Gemma3nTensorNames.inputLayernorm(0)]!! + assertTrue(norm.data is FloatArrayTensorData<*>, "F32 norms stay dense FP32") + } + } + + // --- fixture --------------------------------------------------------------- + + private fun withFixture(block: (indexPath: String) -> Unit) { + val dir = Files.createTempDirectory("gemma3n-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.per_layer_token_embd.weight", listOf(vocab, pleDim), 0.5f, names1) + bf16("model.language_model.per_layer_model_proj.weight", listOf(pleDim, hidden), 0.25f, names1) + f32("model.language_model.per_layer_proj_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) + 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) + bf16("$l.per_layer_projection.weight", listOf(hidden, pleDim), 1.0f, 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) + f32("$l.post_attention_norm.weight", listOf(hidden), 2.5f, names2) + f32("$l.post_ffw_norm.weight", listOf(hidden), 3.0f, names2) + f32("$l.post_norm.weight", listOf(hidden), 3.5f, names2) + bf16("$l.inp_gate.weight", listOf(pleDim, hidden), 0.875f, names2) + bf16("$l.proj.weight", listOf(hidden, pleDim), 1.25f, names2) + // Unmapped decoy: an INT64 tensor the family never asks for. + 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": "gemma3n", + "architectures": ["Gemma3nForConditionalGeneration"], + "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, + "num_kv_shared_layers": 0, + "layer_types": ["full_attention"] + } + } + """.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 + } +}