From 6686107151f629e8f6c817f88775d04bc22fdc18 Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Thu, 3 Sep 2026 22:39:43 +0200 Subject: [PATCH] =?UTF-8?q?functiongemma=20export:=20write=20externals=20f?= =?UTF-8?q?rom=20any=20BufferHandle=20(Owned=20or=20Floats)=20=E2=80=94=20?= =?UTF-8?q?fixes=20the=20ClassCastException=20on=20the=200.53.0=20engine?= =?UTF-8?q?=20loader=20(#405)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Fable 5.1 --- .../FunctionGemmaExportHarness.kt | 95 ++++++++++++++----- 1 file changed, 72 insertions(+), 23 deletions(-) diff --git a/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportHarness.kt b/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportHarness.kt index 00be0485..e48731c4 100644 --- a/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportHarness.kt +++ b/llm-inference/functiongemma/src/jvmMain/kotlin/sk/ainet/models/functiongemma/FunctionGemmaExportHarness.kt @@ -459,28 +459,79 @@ public object FunctionGemmaExportHarness { os.write(ByteBuffer.allocate(8).order(ByteOrder.LITTLE_ENDIAN).putLong(headerBytes.size.toLong()).array()) os.write(headerBytes) for (e in ext) { - val src = e.source as BufferHandle.Owned - if (bf16) { - val data = src.data - val base = src.offset - val n = src.sizeInBytes.toInt() / 4 - val obuf = ByteArray(n * 2) - for (j in 0 until n) { - val o = base + 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)) // truncation = core parity - obuf[j * 2] = (bf and 0xFF).toByte() - obuf[j * 2 + 1] = ((bf ushr 8) and 0xFF).toByte() - } - os.write(obuf) - } else { - os.write(src.data, src.offset, src.sizeInBytes.toInt()) + writeHandle(os, e.source, bf16) + } + } + } + + /** + * Streams one external's f32 payload as raw f32 or truncating bf16 (= core parity), in 1 MiB + * chunks. Dispatches on the [BufferHandle] subtype: the 0.53.0 engine loader delivers large + * constants (the 262144x640 tied embedding) as [BufferHandle.Floats], not [BufferHandle.Owned] — + * the cast that #396 removed from the Gemma 3n harness was still here (#405). + */ + private fun writeHandle(os: java.io.OutputStream, src: BufferHandle, bf16: Boolean) { + val n: Int + val floatAt: (Int) -> Float + when (src) { + is BufferHandle.Owned -> { + val data = src.data + val base = src.offset + n = (src.sizeInBytes / 4).toInt() + if (!bf16) { + os.write(data, base, n * 4) + return + } + floatAt = { j -> + val o = base + j * 4 + Float.fromBits( + (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), + ) } } + is BufferHandle.Floats -> { + val f = src.data + n = f.size + floatAt = { j -> f[j] } + } + else -> error("unsupported BufferHandle ${src::class.simpleName} (${src.sizeInBytes} B)") } + val bpe = if (bf16) 2 else 4 + val chunk = (1 shl 20) / bpe + val buf = ByteArray(chunk * bpe) + var j = 0 + while (j < n) { + val m = minOf(chunk, n - j) + if (bf16) { + for (k in 0 until m) { + val bf = Bf16TensorData.floatToBf16Bits(floatAt(j + k)) // truncation = core parity + buf[k * 2] = (bf and 0xFF).toByte() + buf[k * 2 + 1] = ((bf ushr 8) and 0xFF).toByte() + } + } else { + val bb = ByteBuffer.wrap(buf).order(ByteOrder.LITTLE_ENDIAN) + for (k in 0 until m) bb.putFloat(floatAt(j + k)) + } + os.write(buf, 0, m * bpe) + j += m + } + } + + /** + * Little-endian f32 bytes + base offset for any [BufferHandle]; a [BufferHandle.Floats] is + * materialised once (the quantizer reads rows by byte offset). See [writeHandle] / #405. + */ + private fun ownedBytes(src: BufferHandle): Pair = when (src) { + is BufferHandle.Owned -> src.data to src.offset + is BufferHandle.Floats -> { + val f = src.data + val b = ByteArray(f.size * 4) + val bb = ByteBuffer.wrap(b).order(ByteOrder.LITTLE_ENDIAN) + for (x in f) bb.putFloat(x) + b to 0 + } + else -> error("unsupported BufferHandle ${src::class.simpleName} (${src.sizeInBytes} B)") } private fun voidF32(shape: Shape): sk.ainet.lang.tensor.Tensor = @@ -614,9 +665,7 @@ public object FunctionGemmaExportHarness { os.write(ByteBuffer.allocate(8).order(ByteOrder.LITTLE_ENDIAN).putLong(headerBytes.size.toLong()).array()) os.write(headerBytes) for (e in ext) { - val src = e.source as BufferHandle.Owned - val data = src.data - val base = src.offset + val (data, base) = ownedBytes(e.source) val qs = quant[e.key] if (qs != null) { val (rows, cols) = qs @@ -638,7 +687,7 @@ public object FunctionGemmaExportHarness { os.write(q) os.write(sb) } else { - val n = src.sizeInBytes.toInt() / 4 + val n = e.source.sizeInBytes.toInt() / 4 val ob = ByteArray(n * 2) for (j in 0 until n) { val bf = Bf16TensorData.floatToBf16Bits(leF32(data, base + j * 4))