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
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,13 @@ package sk.ainet.io.model
*
* Shared across all weight loaders (LLaMA, Gemma, etc.).
*/
@Deprecated(
"One of the three axes WeightForm replaces (#1109): this is the encoding axis. No ReplaceWith, " +
"because EncodingRequest is not a drop-in — NATIVE_OPTIMIZED becomes KeepAsStored, " +
"DEQUANTIZE_TO_FP32 becomes DequantizeTo(FP32), and RAW_BYTES has no counterpart because " +
"no loader ever supported it. A wrong ReplaceWith would be worse than none.",
level = DeprecationLevel.WARNING,
)
public enum class QuantPolicy {
/** Keep quantized payloads as raw bytes (Int8 tensor) with quantized shape. */
RAW_BYTES,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,15 @@ package sk.ainet.io.model
* bytes live*. The two are the axes of one loader — `quantPolicy × staging` — instead of the
* separate code paths ("streaming loader" vs "mapped weights helper") they used to be.
*/
@Deprecated(
"One of the three axes WeightForm replaces (#1109): this is the residency axis. Same two " +
"values, resolved from the device's PlannerProfile rather than set by the caller.",
ReplaceWith(
"WeightResidency",
"sk.ainet.lang.memory.plan.WeightResidency",
),
DeprecationLevel.WARNING,
)
public enum class StagingPolicy {
/** Read tensor bytes onto the heap. The historical behaviour, and the only option in a browser. */
HEAP,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,15 @@ package sk.ainet.io.model
* computes the wrong permutation — or refuses, when `out` is not a multiple of the block size. Both
* failures are in the census.
*/
@Deprecated(
"One of the three axes WeightForm replaces (#1109): this is the shape axis. Same two values, " +
"asked as part of one resolved decision instead of a separate flag.",
ReplaceWith(
"WeightShapeOrientation",
"sk.ainet.lang.memory.plan.WeightShapeOrientation",
),
DeprecationLevel.WARNING,
)
public enum class WeightOrientation {
/**
* The file's own order, unreversed — GGUF `ne`, so `[in, out]` for a 2-D weight. What the
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,11 @@ import sk.ainet.context.ExecutionContext
import sk.ainet.io.ParametersLoader
import sk.ainet.io.RandomAccessSource
import sk.ainet.io.gguf.dequant.DequantOps
import sk.ainet.lang.memory.plan.EncodingRequest
import sk.ainet.lang.memory.plan.WeightByteOrder
import sk.ainet.lang.memory.plan.WeightForm
import sk.ainet.lang.memory.plan.WeightResidency
import sk.ainet.lang.memory.plan.WeightShapeOrientation
import sk.ainet.io.model.QuantPolicy
import sk.ainet.io.model.StagingPolicy
import sk.ainet.io.model.WeightOrientation
Expand Down Expand Up @@ -97,14 +102,82 @@ public class StreamingGgufParametersLoader(
* changes what every consumer sees. New code should ask for `OUT_IN`.
*/
private val weightOrientation: WeightOrientation = WeightOrientation.AS_STORED,
/**
* The form weights should take in memory, as one decision instead of three (#1109, #1115).
*
* [quantPolicy], [staging] and [weightOrientation] are the same three axes asked separately,
* and asked of the *caller* — who has to answer for a device they may not be building for.
* `WeightFormResolver.resolve(stored, profile, capabilities)` answers instead, from what the
* file holds, what the device is, and what the backend's kernels can feed.
*
* `null` (the default) means "use the three parameters", so every existing caller is
* byte-identical. Passing a form while also setting any of the three is rejected rather than
* silently resolved, so nobody loses a setting they thought they had.
*/
private val weightForm: WeightForm? = null,
) : ParametersLoader {

/** The three axes as one value: [weightForm] if given, otherwise what the three parameters say. */
private val form: WeightForm = weightForm ?: WeightForm(
encoding = when (quantPolicy) {
QuantPolicy.DEQUANTIZE_TO_FP32 -> EncodingRequest.DequantizeTo(FP32)
else -> EncodingRequest.KeepAsStored
},
order = WeightByteOrder.AS_STORED,
shape = when (weightOrientation) {
WeightOrientation.OUT_IN -> WeightShapeOrientation.OUT_IN
else -> WeightShapeOrientation.AS_STORED
},
residency = when (staging) {
StagingPolicy.MAPPED -> WeightResidency.MAPPED
else -> WeightResidency.HEAP
},
)

/** [QuantPolicy.DEQUANTIZE_TO_FP32] asked for through [form], whichever way it was set. */
private val dequantizeToDense: Boolean = form.encoding is EncodingRequest.DequantizeTo

/** [StagingPolicy.MAPPED] asked for through [form]. */
private val mapsTheFile: Boolean = form.residency == WeightResidency.MAPPED

/** [WeightOrientation.OUT_IN] asked for through [form]. */
private val reversesWeightShape: Boolean = form.shape == WeightShapeOrientation.OUT_IN

init {
require(quantPolicy != QuantPolicy.RAW_BYTES) {
"StreamingGgufParametersLoader does not support QuantPolicy.RAW_BYTES — quantized " +
"tensors are preserved as packed block TensorData (NATIVE_OPTIMIZED) or " +
"dequantized to dense FP32 (DEQUANTIZE_TO_FP32)."
}
if (weightForm != null) {
require(
quantPolicy == QuantPolicy.NATIVE_OPTIMIZED &&
staging == StagingPolicy.HEAP &&
weightOrientation == WeightOrientation.AS_STORED,
) {
"a WeightForm and the quantPolicy/staging/weightOrientation parameters were both " +
"set. They are the same three axes, so one of them would have been silently " +
"ignored; pass only the form."
}
}
require(form.order == WeightByteOrder.AS_STORED) {
"WeightByteOrder.KERNEL_FEED is not supported by this loader yet (#1120). The bytes are " +
"easy — TensorView.prepack(INPUT_BLOCK_MAJOR) does the permutation — but packed " +
"TensorData addresses packedData as canonical row-major in getBlockScale, getCode " +
"and dequantizeBlock, so feed-order bytes would decode the wrong elements without " +
"failing (#973, #968). Packed storage has to be able to declare its own order first."
}
val requested = form.encoding
require(requested !is EncodingRequest.RequantizeTo) {
"EncodingRequest.RequantizeTo(${requested.let { (it as? EncodingRequest.RequantizeTo)?.encoding?.name }}) " +
"is not supported by this loader: re-quantizing a weight the file does not already " +
"carry needs a quantizer per target encoding, and none of them exist here yet."
}
val dequantTarget = (requested as? EncodingRequest.DequantizeTo)?.dtype
require(dequantTarget == null || dequantTarget == FP32) {
"EncodingRequest.DequantizeTo(${dequantTarget?.name}) is not supported: this loader " +
"dequantizes to FP32 only."
}
}

/**
Expand All @@ -114,7 +187,7 @@ public class StreamingGgufParametersLoader(
*/
private fun shapeOf(tensorInfo: StreamingTensorInfo): Shape {
val dims = tensorInfo.shape.map { it.toInt() }
val ordered = if (weightOrientation == WeightOrientation.OUT_IN && dims.size == 2) dims.reversed() else dims
val ordered = if (reversesWeightShape && dims.size == 2) dims.reversed() else dims
return Shape(*ordered.toIntArray())
}

Expand All @@ -127,7 +200,7 @@ public class StreamingGgufParametersLoader(
val source = sourceProvider()
// MAPPED staging needs a file to map; a Blob or an in-memory source has no path and
// silently stays on the heap, which is the documented fallback rather than a failure.
val mapped = if (staging == StagingPolicy.MAPPED) source.filePath?.let { openMappedFile(it) } else null
val mapped = if (mapsTheFile) source.filePath?.let { openMappedFile(it) } else null
try {
StreamingGGUFReader.open(source).use { reader ->
val tensors = reader.tensors
Expand Down Expand Up @@ -254,7 +327,7 @@ public class StreamingGgufParametersLoader(
tensorInfo: StreamingTensorInfo,
rawBytes: ByteArray,
): Tensor<T, V> {
if (quantPolicy == QuantPolicy.DEQUANTIZE_TO_FP32 &&
if (dequantizeToDense &&
(dtype == FP32::class || dtype == FP16::class)
) {
val dest = DequantOps.dequantFromBytes(rawBytes, tensorInfo.tensorType, tensorInfo.nElements.toInt())
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,179 @@
package sk.ainet.io.gguf

import kotlinx.coroutines.runBlocking
import sk.ainet.context.DefaultDataExecutionContext
import sk.ainet.io.JvmRandomAccessSource
import sk.ainet.io.model.QuantPolicy
import sk.ainet.io.model.StagingPolicy
import sk.ainet.io.model.WeightOrientation
import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.plan.EncodingRequest
import sk.ainet.lang.memory.plan.WeightByteOrder
import sk.ainet.lang.memory.plan.WeightForm
import sk.ainet.lang.memory.plan.WeightResidency
import sk.ainet.lang.memory.plan.WeightShapeOrientation
import sk.ainet.lang.tensor.Tensor
import sk.ainet.lang.types.FP32
import java.io.File
import kotlin.test.Test
import kotlin.test.assertContentEquals
import kotlin.test.assertEquals
import kotlin.test.assertFailsWith
import kotlin.test.assertTrue

/**
* #1115: the loader takes one `WeightForm` where it took three flags, and the change is a change of
* *spelling only*.
*
* That is the claim worth testing, because it is the one that can quietly be false. Every
* combination of the three deprecated parameters is loaded twice — once through them, once through
* the `WeightForm` they map to — and the two must agree on shapes and on every element. If the
* mapping is wrong anywhere, some cell of that product disagrees.
*/
@Suppress("DEPRECATION")
@OptIn(ExperimentalMemoryApi::class)
class WeightFormLoaderParityTest {

/** Mixed encodings, and a 2-D weight so the shape axis has something to reverse. */
private fun file(): File = SyntheticGguf.write(
SyntheticGguf.tensor("w_f32", GGMLQuantizationType.F32, elements = 1024),
SyntheticGguf.tensor("w_q4k", GGMLQuantizationType.Q4_K, elements = 1024),
SyntheticGguf.tensor("w_q80", GGMLQuantizationType.Q8_0, elements = 768)
.copy(dims = listOf(256L, 3L)),
SyntheticGguf.tensor("w_f16", GGMLQuantizationType.F16, elements = 1024),
)

private fun loadVia(f: File, build: (() -> JvmRandomAccessSource) -> StreamingGgufParametersLoader):
Map<String, Tensor<FP32, Float>> {
val ctx = DefaultDataExecutionContext()
val loaded = LinkedHashMap<String, Tensor<FP32, Float>>()
runBlocking {
build { JvmRandomAccessSource.open(f) }
.load<FP32, Float>(ctx, FP32::class) { name, tensor -> loaded[name] = tensor }
}
return loaded
}

@Test
fun `every combination of the three flags loads identically through the form it maps to`() {
val f = file()
try {
for (quant in listOf(QuantPolicy.NATIVE_OPTIMIZED, QuantPolicy.DEQUANTIZE_TO_FP32)) {
for (staging in listOf(StagingPolicy.HEAP, StagingPolicy.MAPPED)) {
for (orientation in listOf(WeightOrientation.AS_STORED, WeightOrientation.OUT_IN)) {
val viaFlags = loadVia(f) { src ->
StreamingGgufParametersLoader(
sourceProvider = src,
quantPolicy = quant,
staging = staging,
weightOrientation = orientation,
)
}
val viaForm = loadVia(f) { src ->
StreamingGgufParametersLoader(
sourceProvider = src,
weightForm = WeightForm(
encoding = when (quant) {
QuantPolicy.DEQUANTIZE_TO_FP32 -> EncodingRequest.DequantizeTo(FP32)
else -> EncodingRequest.KeepAsStored
},
shape = when (orientation) {
WeightOrientation.OUT_IN -> WeightShapeOrientation.OUT_IN
else -> WeightShapeOrientation.AS_STORED
},
residency = when (staging) {
StagingPolicy.MAPPED -> WeightResidency.MAPPED
else -> WeightResidency.HEAP
},
),
)
}

val label = "$quant/$staging/$orientation"
assertEquals(viaFlags.keys, viaForm.keys, "$label: different tensors came out")
for ((name, flagsTensor) in viaFlags) {
val formTensor = viaForm.getValue(name)
assertEquals(flagsTensor.shape, formTensor.shape, "$label: $name shape")
assertContentEquals(
flagsTensor.data.copyToFloatArray(),
formTensor.data.copyToFloatArray(),
"$label: $name values",
)
}
}
}
}
} finally {
f.delete()
}
}

@Test
fun `the default loader is the default form`() {
val f = file()
try {
val implicit = loadVia(f) { src -> StreamingGgufParametersLoader(sourceProvider = src) }
val explicit = loadVia(f) { src ->
StreamingGgufParametersLoader(sourceProvider = src, weightForm = WeightForm.AS_STORED_ON_HEAP)
}
for ((name, tensor) in implicit) {
assertContentEquals(
tensor.data.copyToFloatArray(), explicit.getValue(name).data.copyToFloatArray(), name,
)
}
} finally {
f.delete()
}
}

@Test
fun `setting both a form and a flag is refused rather than silently resolved`() {
// One of the two would have to lose, and a caller who set a flag believes it is in effect.
val failure = assertFailsWith<IllegalArgumentException> {
StreamingGgufParametersLoader(
sourceProvider = { JvmRandomAccessSource.open(file()) },
staging = StagingPolicy.MAPPED,
weightForm = WeightForm.AS_STORED_ON_HEAP,
)
}
assertTrue(failure.message!!.contains("pass only the form"), failure.message!!)
}

@Test
fun `KERNEL_FEED is refused for the reason it is refused`() {
// Not "unsupported": the bytes are the easy part. The refusal is because packed TensorData
// reads packedData as canonical row-major, so feed-order bytes decode wrong silently (#1120).
val failure = assertFailsWith<IllegalArgumentException> {
StreamingGgufParametersLoader(
sourceProvider = { JvmRandomAccessSource.open(file()) },
weightForm = WeightForm(order = WeightByteOrder.KERNEL_FEED),
)
}
val message = failure.message!!
assertTrue(message.contains("#1120"), "it names where this is being fixed: $message")
assertTrue(message.contains("canonical row-major"), "and why it cannot be faked: $message")
}

@Test
fun `an encoding request this loader cannot honour is refused up front`() {
val requantize = assertFailsWith<IllegalArgumentException> {
StreamingGgufParametersLoader(
sourceProvider = { JvmRandomAccessSource.open(file()) },
weightForm = WeightForm(
encoding = EncodingRequest.RequantizeTo(sk.ainet.lang.tensor.storage.TensorEncoding.Q8_0),
),
)
}
assertTrue(requantize.message!!.contains("quantizer"), requantize.message!!)

val toFp16 = assertFailsWith<IllegalArgumentException> {
StreamingGgufParametersLoader(
sourceProvider = { JvmRandomAccessSource.open(file()) },
weightForm = WeightForm(
encoding = EncodingRequest.DequantizeTo(sk.ainet.lang.types.FP16),
),
)
}
assertTrue(toFp16.message!!.contains("FP32 only"), toFp16.message!!)
}
}
Loading
Loading