From fc6582f003956c6ce4dfcb5f9eae67c3a216556b Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Wed, 2 Sep 2026 11:50:39 +0200 Subject: [PATCH] fix(gemma3n): stream export weights through the BufferResolver, drop the Owned cast (SKaiNET#1247) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit writeSafetensors cast every ExternalParameterRef.source to BufferHandle.Owned. With SKaiNET#1247 the engine hands >=2 GiB FP32 constants over as an aliased FloatArray (BufferHandle.Floats) that never exists as one ByteArray — the tied embedding is exactly Int.MAX_VALUE + 1 bytes — so the cast threw ClassCastException the moment the strict conversion finally got that far. Reading through DefaultBufferResolver (present in the published engine) keeps the harness compiling against 0.52.0 and streams any handle in bounded chunks with the same chunked bf16 conversion. Verified: GEMMA3N_LAYERS=4 export against the #1247 engine branches emits gemma3n-gen.mlir + a 1546 MiB safetensors (124 params) — the first time this export has produced a servable module. Co-Authored-By: Claude Fable 5.1 --- .../models/gemma3n/Gemma3nExportHarness.kt | 52 +++++++++++-------- 1 file changed, 30 insertions(+), 22 deletions(-) diff --git a/llm-inference/gemma3n/src/jvmMain/kotlin/sk/ainet/models/gemma3n/Gemma3nExportHarness.kt b/llm-inference/gemma3n/src/jvmMain/kotlin/sk/ainet/models/gemma3n/Gemma3nExportHarness.kt index 1110d552..4ff4321e 100644 --- a/llm-inference/gemma3n/src/jvmMain/kotlin/sk/ainet/models/gemma3n/Gemma3nExportHarness.kt +++ b/llm-inference/gemma3n/src/jvmMain/kotlin/sk/ainet/models/gemma3n/Gemma3nExportHarness.kt @@ -20,7 +20,7 @@ import sk.ainet.lang.tensor.VoidOpsTensor import sk.ainet.lang.tensor.data.Bf16TensorData import sk.ainet.lang.tensor.data.TensorData import sk.ainet.lang.tensor.ops.VoidTensorOps -import sk.ainet.lang.tensor.storage.BufferHandle +import sk.ainet.lang.tensor.storage.DefaultBufferResolver import sk.ainet.lang.types.FP32 import sk.ainet.tape.Execution import java.io.BufferedOutputStream @@ -265,30 +265,38 @@ public object Gemma3nExportHarness { os.write(ByteBuffer.allocate(8).order(ByteOrder.LITTLE_ENDIAN).putLong(headerBytes.size.toLong()).array()) os.write(headerBytes) val chunkElems = 1 shl 24 // 16M floats per conversion chunk (32 MiB bf16 out) + // Read every handle through the resolver instead of casting to + // BufferHandle.Owned: with SKaiNET#1247 the engine hands ≥2 GiB + // FP32 constants over as an aliased FloatArray (BufferHandle.Floats) + // that never exists as one ByteArray, and the resolver streams it + // as little-endian bytes in bounded chunks. + val resolver = DefaultBufferResolver() for (e in ext) { - val src = e.source as BufferHandle.Owned - if (bf16) { - val data = src.data - val n = (src.sizeInBytes / 4).toInt() - var done = 0 - while (done < n) { - val take = minOf(chunkElems, n - done) - val obuf = ByteArray(take * 2) - for (j in 0 until take) { - val o = src.offset + (done + j) * 4 - val fb = (data[o].toInt() and 0xFF) or - ((data[o + 1].toInt() and 0xFF) shl 8) or - ((data[o + 2].toInt() and 0xFF) shl 16) or - ((data[o + 3].toInt() and 0xFF) shl 24) - val bf = Bf16TensorData.floatToBf16Bits(Float.fromBits(fb)) - obuf[j * 2] = (bf and 0xFF).toByte() - obuf[j * 2 + 1] = ((bf ushr 8) and 0xFF).toByte() + resolver.resolve(e.source).use { acc -> + val total = acc.sizeInBytes + var doneBytes = 0L + while (doneBytes < total) { + val takeBytes = minOf(chunkElems.toLong() * 4L, total - doneBytes).toInt() + val data = acc.readBytes(doneBytes, takeBytes) + if (bf16) { + val take = takeBytes / 4 + val obuf = ByteArray(take * 2) + for (j in 0 until take) { + val o = j * 4 + val fb = (data[o].toInt() and 0xFF) or + ((data[o + 1].toInt() and 0xFF) shl 8) or + ((data[o + 2].toInt() and 0xFF) shl 16) or + ((data[o + 3].toInt() and 0xFF) shl 24) + val bf = Bf16TensorData.floatToBf16Bits(Float.fromBits(fb)) + obuf[j * 2] = (bf and 0xFF).toByte() + obuf[j * 2 + 1] = ((bf ushr 8) and 0xFF).toByte() + } + os.write(obuf) + } else { + os.write(data) } - os.write(obuf) - done += take + doneBytes += takeBytes } - } else { - os.write(src.data, src.offset, src.sizeInBytes.toInt()) } } }