Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,18 @@

## [Unreleased]

### Added

- **`BufferHandle.Floats` — array-free path for ≥2 GiB constants**
([#1247](https://github.com/SKaiNET-developers/SKaiNET/issues/1247)): external FP32 constants
now ride the aliased `FloatArray` end-to-end (graph → `ExternalParameterRef` → `.irpa`), never
serializing to a single `ByteArray` — the gemma3n tied embedding (262144x2048 FP32 =
`Int.MAX_VALUE` + 1 bytes) structurally cannot exist as one byte buffer. `IrpaWriter` streams
the values little-endian in 64 MiB chunks; `DefaultBufferResolver` reads through a chunked byte
view. Constant element counts fold in `Long`, and an oversized single-buffer serialization now
throws `ConstantTooLargeException` with the remediation instead of the
`NegativeArraySizeException` that was previously mistaken for a registry miss.

### Changed

- **Graph constants alias live weights; packed params fail loudly**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,10 @@ public final class sk/ainet/compile/hlo/ConstantMaterializationPolicy$SizeThresh
public fun toString ()Ljava/lang/String;
}

public final class sk/ainet/compile/hlo/ConstantTooLargeException : java/lang/IllegalStateException {
public fun <init> (Ljava/lang/String;)V
}

public final class sk/ainet/compile/hlo/ConversionContext {
public fun <init> (Lsk/ainet/compile/hlo/TypeMapper;)V
public fun <init> (Lsk/ainet/compile/hlo/TypeMapper;Lsk/ainet/lang/graph/ComputeGraph;)V
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ internal fun numberListToLittleEndianBytes(
expectedElements: Int
): ByteArray {
val count = expectedElements.coerceAtLeast(values.size)
requireSerializableByteCount(count, bytesPerSerializedElement(dtype), dtype)
val normalized = dtype.uppercase()

return when (normalized) {
Expand Down Expand Up @@ -91,6 +92,7 @@ internal fun floatArrayToLittleEndianBytes(
expectedElements: Int
): ByteArray {
val count = expectedElements.coerceAtLeast(values.size)
requireSerializableByteCount(count, bytesPerSerializedElement(dtype), dtype)
val n = minOf(count, values.size)
return when (dtype.uppercase()) {
"FP32", "F32", "FLOAT32" -> {
Expand Down Expand Up @@ -125,9 +127,60 @@ internal fun floatArrayToLittleEndianBytes(
* Expected element count for a (possibly empty) shape. Empty shape
* (scalar) means one element; `null` / absent dims degrade to 0 so the
* caller can detect "no declared shape".
*
* [Long] arithmetic (#1247): the gemma3n token embedding is
* 262144 x 2048 = 536,870,912 elements — an [Int] fold of its *byte*
* count goes negative, which previously surfaced as a
* `NegativeArraySizeException` inside the serializer and was mistaken
* for a registry miss ("Unsupported op 'weight' … Known names: […]").
*/
internal fun elementCountFromShape(shape: List<Int>?): Long {
if (shape == null) return 0L
if (shape.isEmpty()) return 1L
return shape.fold(1L) { acc, d -> acc * d }
}

/**
* A constant's serialized form would exceed the JVM single-array ceiling.
* Deliberately NOT an [IllegalArgumentException]: the converter's
* "unsupported dtype" fallback catches that type to retry inline emission,
* and inlining a multi-GiB tensor as text is exactly the wrong recovery.
*/
public class ConstantTooLargeException(message: String) : IllegalStateException(message)

private fun bytesPerSerializedElement(dtype: String): Long = when (dtype.uppercase()) {
"FP64", "F64", "FLOAT64", "I64", "INT64" -> 8L
else -> 4L
}

/**
* Narrow a [Long] element count for the byte-serialization paths, which
* address a single array. Throws [ConstantTooLargeException] instead of
* truncating — truncation is how the #1247 embedding turned into a
* `NegativeArraySizeException`.
*/
internal fun elementCountFromShape(shape: List<Int>?): Int {
if (shape == null) return 0
if (shape.isEmpty()) return 1
return shape.fold(1) { acc, d -> acc * d }
internal fun checkedIntElements(elementCount: Long): Int {
if (elementCount > Int.MAX_VALUE - 8L) {
throw ConstantTooLargeException(
"Constant of $elementCount elements exceeds the single-array serialization " +
"ceiling. Use ConstantMaterializationPolicy.ExternalAlways with FP32 values — " +
"the external path carries a FloatArray (BufferHandle.Floats) without byte " +
"serialization (issue #1247)."
)
}
return elementCount.toInt()
}

private fun requireSerializableByteCount(count: Int, bytesPerElement: Long, dtype: String) {
val byteCount = count.toLong() * bytesPerElement
// A JVM array tops out just under Int.MAX_VALUE entries; keep a small
// margin for VM-specific header overhead.
if (byteCount > Int.MAX_VALUE - 8L) {
throw ConstantTooLargeException(
"Constant of $count $dtype elements needs $byteCount bytes; single-buffer " +
"serialization caps at 2 GiB - 1. Use ConstantMaterializationPolicy.ExternalAlways " +
"with FP32 values — the external path carries a FloatArray (BufferHandle.Floats) " +
"without any byte serialization (issue #1247)."
)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import sk.ainet.compile.hlo.ConversionContext
import sk.ainet.compile.hlo.ConversionResult
import sk.ainet.compile.hlo.ExternalParameterRef
import sk.ainet.compile.hlo.StableHloOperationConverter
import sk.ainet.compile.hlo.checkedIntElements
import sk.ainet.compile.hlo.elementCountFromShape
import sk.ainet.compile.hlo.floatArrayToLittleEndianBytes
import sk.ainet.compile.hlo.numberListToLittleEndianBytes
Expand Down Expand Up @@ -410,9 +411,9 @@ public class ConstantOperationsConverter : StableHloOperationConverter {
val encoding = outputSpec.tensorEncoding
?: TensorEncoding.Dense(bytesPerElement = bytesPerElement(outputSpec.dtype))
val elementCount = elementCountFromShape(outputSpec.shape)
if (elementCount <= 0) return null
if (elementCount <= 0L) return null

val logicalBytes = encoding.physicalBytes(elementCount.toLong()) ?: return null
val logicalBytes = encoding.physicalBytes(elementCount) ?: return null
val scope = when (policy) {
is ConstantMaterializationPolicy.InlineAlways -> return null
is ConstantMaterializationPolicy.ExternalAlways -> policy.scope
Expand All @@ -424,9 +425,11 @@ public class ConstantOperationsConverter : StableHloOperationConverter {

// Serialize now. Fall back to inline on unsupported dtype —
// a loud exception here would defeat the "default path is
// safe" invariant of the seam.
// safe" invariant of the seam. A ConstantTooLargeException is NOT
// caught: inlining a multi-GiB tensor as text is the wrong recovery
// (#1247) — it propagates with its actionable message.
val bytes = try {
numberListToLittleEndianBytes(values, outputSpec.dtype, elementCount)
numberListToLittleEndianBytes(values, outputSpec.dtype, checkedIntElements(elementCount))
} catch (e: IllegalArgumentException) {
context.emitComment(
"external materialization fell back to inline for ${node.id}: ${e.message}"
Expand Down Expand Up @@ -489,9 +492,9 @@ public class ConstantOperationsConverter : StableHloOperationConverter {
val encoding = outputSpec.tensorEncoding
?: TensorEncoding.Dense(bytesPerElement = bytesPerElement(outputSpec.dtype))
val elementCount = elementCountFromShape(outputSpec.shape)
if (elementCount <= 0) return null
if (elementCount <= 0L) return null

val logicalBytes = encoding.physicalBytes(elementCount.toLong()) ?: return null
val logicalBytes = encoding.physicalBytes(elementCount) ?: return null
val scope = when (policy) {
is ConstantMaterializationPolicy.InlineAlways -> return null
is ConstantMaterializationPolicy.ExternalAlways -> policy.scope
Expand All @@ -501,13 +504,31 @@ public class ConstantOperationsConverter : StableHloOperationConverter {
}
}

val bytes = try {
floatArrayToLittleEndianBytes(values, outputSpec.dtype, elementCount)
} catch (e: IllegalArgumentException) {
context.emitComment(
"external materialization fell back to inline for ${node.id}: ${e.message}"
)
return null
val normalizedDtype = outputSpec.dtype.uppercase()
val source: BufferHandle = if (
(normalizedDtype == "FP32" || normalizedDtype == "F32" || normalizedDtype == "FLOAT32") &&
values.size.toLong() == elementCount
) {
// Array-free path (#1247): the FloatArray aliases the live module
// weight all the way from TraceToGraphBuilder — hand it to the
// packager as-is. No byte serialization means no extra copy and
// no 2 GiB ByteArray ceiling: the gemma3n tied embedding
// (262144x2048 = Int.MAX_VALUE + 1 bytes) only exports this way.
BufferHandle.Floats(values)
} else {
// Under-filled initializations and non-FP32 dtypes keep the
// padded byte serialization. ConstantTooLargeException is NOT
// caught — inlining a multi-GiB tensor as text is the wrong
// recovery; it propagates with its actionable message.
val bytes = try {
floatArrayToLittleEndianBytes(values, outputSpec.dtype, checkedIntElements(elementCount))
} catch (e: IllegalArgumentException) {
context.emitComment(
"external materialization fell back to inline for ${node.id}: ${e.message}"
)
return null
}
BufferHandle.Owned(bytes)
}

val key = outputSpec.name.ifEmpty { node.id }
Expand All @@ -516,7 +537,7 @@ public class ConstantOperationsConverter : StableHloOperationConverter {
scope = scope,
key = key,
encoding = encoding,
source = BufferHandle.Owned(bytes),
source = source,
blockOrder = outputSpec.blockOrder,
)
)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,148 @@
package sk.ainet.compile.hlo

import sk.ainet.lang.graph.DefaultComputeGraph
import sk.ainet.lang.graph.GraphEdge
import sk.ainet.lang.graph.GraphNode
import sk.ainet.lang.tensor.Tensor
import sk.ainet.lang.tensor.ops.Operation
import sk.ainet.lang.tensor.ops.TensorSpec
import sk.ainet.lang.tensor.ops.ValidationResult
import sk.ainet.lang.tensor.storage.BufferHandle
import sk.ainet.lang.types.DType
import kotlin.test.Test
import kotlin.test.assertEquals
import kotlin.test.assertFailsWith
import kotlin.test.assertSame
import kotlin.test.assertTrue

/**
* #1247 big-constant contract: element counts fold in Long (the gemma3n
* embedding is 262144 x 2048 = Int.MAX_VALUE + 1 BYTES — an Int fold went
* negative and surfaced as a NegativeArraySizeException mistaken for a
* registry miss), oversized single-buffer serialization refuses with an
* actionable message instead of throwing array-size garbage, and the
* external FP32 path hands the aliased FloatArray to the packager as
* [BufferHandle.Floats] with no byte serialization at all.
*/
class BigConstantMaterializationTest {

@Test
fun elementCountFromShape_folds_in_long() {
assertEquals(536_870_912L, elementCountFromShape(listOf(262_144, 2_048)))
assertEquals(1L, elementCountFromShape(emptyList()))
assertEquals(0L, elementCountFromShape(null))
// 2 Gi elements — Int fold would be 0/negative; Long must be exact.
assertEquals(2_147_483_648L, elementCountFromShape(listOf(65_536, 32_768)))
}

@Test
fun checkedIntElements_refuses_oversized_counts_with_actionable_message() {
val e = assertFailsWith<ConstantTooLargeException> {
checkedIntElements(536_870_912L * 4L)
}
assertTrue("ExternalAlways" in (e.message ?: ""), "refusal must point at the external path")
assertEquals(1024, checkedIntElements(1024L))
}

@Test
fun serializer_refuses_byte_overflow_instead_of_negative_array_size() {
// 536,870,912 FP32 elements = 2 GiB of bytes: previously
// ByteArray(count * 4) threw NegativeArraySizeException.
val e = assertFailsWith<ConstantTooLargeException> {
floatArrayToLittleEndianBytes(FloatArray(0), "FP32", 536_870_912)
}
assertTrue("2 GiB" in (e.message ?: ""), "refusal must state the ceiling: ${e.message}")
}

@Test
fun external_fp32_path_aliases_the_float_array_without_serialization() {
val weights = FloatArray(12) { it * 0.25f }
val module = convertWeightGraph(weights, shape = listOf(4, 3))

val ref = module.externalParameters.single()
val source = ref.source
assertTrue(source is BufferHandle.Floats, "FP32 external constant must ride BufferHandle.Floats, got $source")
assertSame(weights, source.data, "the handle must alias the input array — zero copies end-to-end")
assertEquals(48L, source.sizeInBytes)
assertTrue(module.content.contains("util.global.load"), "external constant must load from a util.global")
}

@Test
fun tied_weight_feeding_two_consumers_emits_one_global() {
val weights = FloatArray(8) { it.toFloat() }
val graph = DefaultComputeGraph()
val weightNode = GraphNode(
id = "w1",
operation = weightOp(weights),
inputs = emptyList(),
outputs = listOf(TensorSpec("tied_embed", listOf(2, 4), "FP32"))
)
val consumerA = computeNode("relu1", "relu", inputs = 1)
val consumerB = computeNode("relu2", "relu", inputs = 1)
graph.addNode(weightNode)
graph.addNode(consumerA)
graph.addNode(consumerB)
graph.addEdge(GraphEdge("ea", weightNode, consumerA, 0, 0, weightNode.outputs[0]))
graph.addEdge(GraphEdge("eb", weightNode, consumerB, 0, 0, weightNode.outputs[0]))

val converter = StableHloConverterFactory.createExtended(
policy = ConstantMaterializationPolicy.ExternalAlways(scope = "model")
)
val module = converter.convert(graph, "tied")

assertEquals(1, module.externalParameters.size, "one tied weight must register exactly one external ref")
val globalDecls = module.content.lines().count { "util.global private @tied_embed" in it }
assertEquals(1, globalDecls, "one tied weight must declare exactly one util.global:\n${module.content}")
}

private fun convertWeightGraph(weights: FloatArray, shape: List<Int>): StableHloModule {
val graph = DefaultComputeGraph()
val weightNode = GraphNode(
id = "w1",
operation = weightOp(weights),
inputs = emptyList(),
outputs = listOf(TensorSpec("w1_spec", shape, "FP32"))
)
val consumer = computeNode("relu1", "relu", inputs = 1)
graph.addNode(weightNode)
graph.addNode(consumer)
graph.addEdge(GraphEdge("e1", weightNode, consumer, 0, 0, weightNode.outputs[0]))

val converter = StableHloConverterFactory.createExtended(
policy = ConstantMaterializationPolicy.ExternalAlways(scope = "model")
)
return converter.convert(graph, "weights")
}

private fun weightOp(values: FloatArray): Operation = object : Operation {
override val name: String = "weight"
override val type: String = "constant"
override val parameters: Map<String, Any> = mapOf(
"initial_value" to values,
"trainable" to false
)
override fun <T : DType, V> execute(inputs: List<Tensor<T, V>>): List<Tensor<T, V>> =
throw UnsupportedOperationException("test fixture only")
override fun validateInputs(inputs: List<TensorSpec>): ValidationResult = ValidationResult.Valid
override fun inferOutputs(inputs: List<TensorSpec>): List<TensorSpec> = emptyList()
override fun clone(newParameters: Map<String, Any>): Operation = this
override fun serialize(): Map<String, Any> = mapOf("name" to name, "type" to type)
}

private fun computeNode(id: String, opName: String, inputs: Int): GraphNode = GraphNode(
id = id,
operation = object : Operation {
override val name: String = opName
override val type: String = "compute"
override val parameters: Map<String, Any> = emptyMap()
override fun <T : DType, V> execute(inputs: List<Tensor<T, V>>): List<Tensor<T, V>> =
throw UnsupportedOperationException("test fixture only")
override fun validateInputs(inputs: List<TensorSpec>): ValidationResult = ValidationResult.Valid
override fun inferOutputs(inputs: List<TensorSpec>): List<TensorSpec> = emptyList()
override fun clone(newParameters: Map<String, Any>): Operation = this
override fun serialize(): Map<String, Any> = mapOf("name" to name, "type" to type)
},
inputs = List(inputs) { TensorSpec("in$it", listOf(2, 4), "FP32") },
outputs = listOf(TensorSpec("$id-out", listOf(2, 4), "FP32"))
)
}
Original file line number Diff line number Diff line change
Expand Up @@ -224,6 +224,7 @@ public class IrpaWriter {
private fun writeBufferHandle(sink: Sink, handle: BufferHandle) {
when (handle) {
is BufferHandle.Owned -> writeByteArray(sink, handle.data, handle.offset, handle.sizeInBytes.toInt())
is BufferHandle.Floats -> writeFloatArrayLe(sink, handle.data)
is BufferHandle.Borrowed -> writeByteArray(sink, handle.data, handle.offset, handle.sizeInBytes.toInt())
is BufferHandle.FileBacked -> writeFileBackedBytes(sink, handle)
else -> throw IllegalArgumentException(
Expand All @@ -235,6 +236,29 @@ public class IrpaWriter {
}
}

private fun writeFloatArrayLe(sink: Sink, data: FloatArray) {
// Stream little-endian in 64 MiB chunks: a Floats handle exists
// precisely because its logical bytes can exceed what a single
// ByteArray can hold (issue #1247, the 2 GiB tied embedding) —
// never materialize the whole payload.
val chunkFloats = 16 * 1024 * 1024
val buf = ByteArray(minOf(chunkFloats, data.size.coerceAtLeast(1)) * 4)
var i = 0
while (i < data.size) {
val n = minOf(chunkFloats, data.size - i)
var b = 0
for (j in i until i + n) {
val bits = data[j].toRawBits()
buf[b++] = (bits and 0xff).toByte()
buf[b++] = (bits ushr 8 and 0xff).toByte()
buf[b++] = (bits ushr 16 and 0xff).toByte()
buf[b++] = (bits ushr 24 and 0xff).toByte()
}
sink.write(buf, 0, b)
i += n
}
}

private fun writeByteArray(sink: Sink, data: ByteArray, offset: Int, length: Int) {
// Bulk range write. Owned/Borrowed buffers can be large for real
// LLM weights (a 262153x640 FP32 embedding = ~670MB); a
Expand Down
Loading
Loading