diff --git a/skainet-apps/skainet-plan/src/main/kotlin/sk/ainet/apps/plan/SkainetPlan.kt b/skainet-apps/skainet-plan/src/main/kotlin/sk/ainet/apps/plan/SkainetPlan.kt index 76a267660..766286361 100644 --- a/skainet-apps/skainet-plan/src/main/kotlin/sk/ainet/apps/plan/SkainetPlan.kt +++ b/skainet-apps/skainet-plan/src/main/kotlin/sk/ainet/apps/plan/SkainetPlan.kt @@ -14,7 +14,9 @@ import sk.ainet.lang.memory.plan.KvCacheMode import sk.ainet.lang.memory.plan.MemoryPlan import sk.ainet.lang.memory.plan.MemoryPlans import sk.ainet.lang.memory.plan.PlanInput +import sk.ainet.lang.memory.plan.KernelCapabilities import sk.ainet.lang.memory.plan.PlannerProfile +import sk.ainet.lang.memory.plan.resolveWeightForms import sk.ainet.lang.memory.plan.ProfiledPlan import java.io.File import kotlin.system.exitProcess @@ -47,6 +49,13 @@ public fun main(args: Array) { fullName = "profile", description = "Device profile whose rules the plan follows (M2-F6): mobile = 2 GB phone, desktop, native", ).default("none") + val kernels by parser.option( + ArgType.Choice(listOf("all", "dense"), { it }), + fullName = "kernels", + description = "What the target's kernels can feed (#1116): all = packed kernels for every " + + "shipped encoding (the default, and what SKaiNET's CPU backend carries); dense = FP32 " + + "only, so the plan prices dequantizing every quantized weight at load", + ).default("all") parser.parse(args) val file = File(model) @@ -56,7 +65,16 @@ public fun main(args: Array) { val profile = profileFor(profileName) val profiled: ProfiledPlan = JvmRandomAccessSource.open(file).use { src -> val reader = StreamingGGUFReader.open(src) - val input = reader.planInput(ctx = ctx, prefillChunk = prefill, kvMode = kvMode) + val stored = reader.planInput(ctx = ctx, prefillChunk = prefill, kvMode = kvMode) + // Which encodings the *target* can feed is not a property of this machine, so it is asked + // for rather than detected. `all` keeps the plan exactly as it was before #1116. + val input = when { + profile == null -> stored + else -> stored.resolveWeightForms( + profile, + if (kernels == "dense") KernelCapabilities.DENSE_ONLY else KernelCapabilities.EVERYTHING, + ) + } val available = budget?.let { parseBytes(it) } ?: Runtime.getRuntime().maxMemory() if (profile != null) { // A profile owns the reserve, so --budget names what the *device* has, not what the diff --git a/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api b/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api index 0c2a45f2d..358b5441f 100644 --- a/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api +++ b/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api @@ -1530,7 +1530,8 @@ public final class sk/ainet/lang/memory/plan/KvCacheMode : java/lang/Enum { } public final class sk/ainet/lang/memory/plan/MemoryPlan { - public fun (Lsk/ainet/lang/memory/plan/PlanInput;JJJJJLsk/ainet/lang/memory/plan/Budget;)V + public fun (Lsk/ainet/lang/memory/plan/PlanInput;JJJJJLsk/ainet/lang/memory/plan/Budget;J)V + public synthetic fun (Lsk/ainet/lang/memory/plan/PlanInput;JJJJJLsk/ainet/lang/memory/plan/Budget;JILkotlin/jvm/internal/DefaultConstructorMarker;)V public final fun component1 ()Lsk/ainet/lang/memory/plan/PlanInput; public final fun component2 ()J public final fun component3 ()J @@ -1538,11 +1539,13 @@ public final class sk/ainet/lang/memory/plan/MemoryPlan { public final fun component5 ()J public final fun component6 ()J public final fun component7 ()Lsk/ainet/lang/memory/plan/Budget; - public final fun copy (Lsk/ainet/lang/memory/plan/PlanInput;JJJJJLsk/ainet/lang/memory/plan/Budget;)Lsk/ainet/lang/memory/plan/MemoryPlan; - public static synthetic fun copy$default (Lsk/ainet/lang/memory/plan/MemoryPlan;Lsk/ainet/lang/memory/plan/PlanInput;JJJJJLsk/ainet/lang/memory/plan/Budget;ILjava/lang/Object;)Lsk/ainet/lang/memory/plan/MemoryPlan; + public final fun component8 ()J + public final fun copy (Lsk/ainet/lang/memory/plan/PlanInput;JJJJJLsk/ainet/lang/memory/plan/Budget;J)Lsk/ainet/lang/memory/plan/MemoryPlan; + public static synthetic fun copy$default (Lsk/ainet/lang/memory/plan/MemoryPlan;Lsk/ainet/lang/memory/plan/PlanInput;JJJJJLsk/ainet/lang/memory/plan/Budget;JILjava/lang/Object;)Lsk/ainet/lang/memory/plan/MemoryPlan; public fun equals (Ljava/lang/Object;)Z public final fun getBudget ()Lsk/ainet/lang/memory/plan/Budget; public final fun getFits ()Ljava/lang/Boolean; + public final fun getFormConversionBytes ()J public final fun getForwardBytes ()J public final fun getHeadroomBytes ()J public final fun getInput ()Lsk/ainet/lang/memory/plan/PlanInput; @@ -1551,6 +1554,7 @@ public final class sk/ainet/lang/memory/plan/MemoryPlan { public final fun getLines ()Ljava/util/List; public final fun getResidentBytes ()J public final fun getTotalBytes ()J + public final fun getWeightsAsStoredBytes ()J public final fun getWeightsBytes ()J public fun hashCode ()I public final fun render ()Ljava/lang/String; @@ -1648,21 +1652,25 @@ public final class sk/ainet/lang/memory/plan/PlanLine { } public final class sk/ainet/lang/memory/plan/PlanTensor { - public fun (Ljava/lang/String;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/memory/Format;JJ)V + public fun (Ljava/lang/String;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/memory/Format;JJLsk/ainet/lang/memory/plan/WeightForm;)V + public synthetic fun (Ljava/lang/String;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/memory/Format;JJLsk/ainet/lang/memory/plan/WeightForm;ILkotlin/jvm/internal/DefaultConstructorMarker;)V public final fun component1 ()Ljava/lang/String; public final fun component2 ()Lsk/ainet/lang/tensor/TensorId; public final fun component3 ()Lsk/ainet/lang/memory/Format; public final fun component4 ()J public final fun component5 ()J - public final fun copy (Ljava/lang/String;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/memory/Format;JJ)Lsk/ainet/lang/memory/plan/PlanTensor; - public static synthetic fun copy$default (Lsk/ainet/lang/memory/plan/PlanTensor;Ljava/lang/String;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/memory/Format;JJILjava/lang/Object;)Lsk/ainet/lang/memory/plan/PlanTensor; + public final fun component6 ()Lsk/ainet/lang/memory/plan/WeightForm; + public final fun copy (Ljava/lang/String;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/memory/Format;JJLsk/ainet/lang/memory/plan/WeightForm;)Lsk/ainet/lang/memory/plan/PlanTensor; + public static synthetic fun copy$default (Lsk/ainet/lang/memory/plan/PlanTensor;Ljava/lang/String;Lsk/ainet/lang/tensor/TensorId;Lsk/ainet/lang/memory/Format;JJLsk/ainet/lang/memory/plan/WeightForm;ILjava/lang/Object;)Lsk/ainet/lang/memory/plan/PlanTensor; public fun equals (Ljava/lang/Object;)Z public final fun getAllocation ()Lsk/ainet/lang/memory/AllocationSpec; public final fun getBytes ()J public final fun getElementCount ()J + public final fun getForm ()Lsk/ainet/lang/memory/plan/WeightForm; public final fun getFormat ()Lsk/ainet/lang/memory/Format; public final fun getId ()Lsk/ainet/lang/tensor/TensorId; public final fun getName ()Ljava/lang/String; + public final fun getResidentBytes ()J public fun hashCode ()I public fun toString ()Ljava/lang/String; } @@ -1841,6 +1849,10 @@ public final class sk/ainet/lang/memory/plan/WeightFormResolver { public final fun resolve (Lsk/ainet/lang/tensor/storage/TensorEncoding;Lsk/ainet/lang/memory/plan/PlannerProfile;Lsk/ainet/lang/memory/plan/KernelCapabilities;)Lsk/ainet/lang/memory/plan/WeightForm; } +public final class sk/ainet/lang/memory/plan/WeightFormResolverKt { + public static final fun resolveWeightForms (Lsk/ainet/lang/memory/plan/PlanInput;Lsk/ainet/lang/memory/plan/PlannerProfile;Lsk/ainet/lang/memory/plan/KernelCapabilities;)Lsk/ainet/lang/memory/plan/PlanInput; +} + public final class sk/ainet/lang/memory/plan/WeightResidency : java/lang/Enum { public static final field HEAP Lsk/ainet/lang/memory/plan/WeightResidency; public static final field MAPPED Lsk/ainet/lang/memory/plan/WeightResidency; diff --git a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/memory/plan/MemoryPlan.kt b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/memory/plan/MemoryPlan.kt index f5511e06e..9c84165f9 100644 --- a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/memory/plan/MemoryPlan.kt +++ b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/memory/plan/MemoryPlan.kt @@ -20,7 +20,30 @@ public data class PlanTensor( val elementCount: Long, /** Physical bytes; falls back to the checkpoint's own byte count when the encoding cannot compute it. */ val bytes: Long, + /** + * The form this weight was resolved to take in memory (#1109/#1116), or `null` when nothing has + * resolved one and the file's own bytes are what will be held. + */ + val form: WeightForm? = null, ) { + /** + * What this weight actually costs in memory once [form] is honoured. + * + * [bytes] is what the *file* holds. They differ exactly when the resolved form re-encodes: + * a Q4_K tensor kept as Q4_K costs its packed bytes, and the same tensor dequantized to FP32 + * costs roughly eight times that. Planning the first number while loading the second is how a + * dequantization becomes an OOM instead of a line in a table. + */ + val residentBytes: Long + get() = when (val request = form?.encoding) { + null, EncodingRequest.KeepAsStored -> bytes + is EncodingRequest.DequantizeTo -> + Format.dense(request.dtype).physicalBytes(elementCount) + ?: (request.dtype.sizeInBytes.toLong() * elementCount) + is EncodingRequest.RequantizeTo -> + Format(format.dtype, request.encoding).physicalBytes(elementCount) ?: bytes + } + /** The allocation this weight needs: mapped, model-lifetime, read-only. */ val allocation: AllocationSpec get() = AllocationSpec(format, elementCount, MemoryDomain.MMAP_FILE, ScopeKind.MODEL, mutable = false) @@ -139,7 +162,18 @@ public data class MemoryPlan( val forwardBytes: Long, val headroomBytes: Long, val budget: Budget?, + /** + * What the weights occupy *in the file*, before any resolved [WeightForm] re-encodes them + * (#1116). Equal to [weightsBytes] when nothing was resolved, which is every plan built before + * forms existed — hence the default. + */ + val weightsAsStoredBytes: Long = weightsBytes, ) { + /** + * Bytes the resolved forms add to the weights — a dequantization's price, made visible before + * it is paid rather than discovered as an OOM. Zero when the weights are held as stored. + */ + val formConversionBytes: Long get() = weightsBytes - weightsAsStoredBytes val totalBytes: Long get() = weightsBytes + kvBytes + forwardBytes + headroomBytes val residentBytes: Long get() = weightsBytes + kvBytes @@ -148,7 +182,13 @@ public data class MemoryPlan( val lines: List get() = listOf( - PlanLine("weights", "Mapped, packed", weightsBytes, resident = true), + PlanLine( + "weights", + if (formConversionBytes == 0L) "Mapped, packed" + else "re-encoded at load (+${MemoryPlans.formatBytes(formConversionBytes)})", + weightsBytes, + resident = true, + ), PlanLine("kv cache", input.kvMode.label + " @ ctx ${input.ctx}", kvBytes, resident = true), PlanLine("forward", "prefill chunk ${input.prefillChunk}", forwardBytes, resident = false), PlanLine("heap", "headroom", headroomBytes, resident = false), @@ -167,6 +207,14 @@ public data class MemoryPlan( val half = MemoryPlans.plan(input.copy(ctx = halfCtx), budget) out += Suggestion("--ctx $halfCtx", totalBytes - half.totalBytes) } + if (formConversionBytes > 0) { + // Worth saying first: unlike ctx or KV mode, this cost was not asked for — it is what + // the resolver chose because no kernel on the target could feed the stored encoding. + out += Suggestion( + "a build with kernels for the stored encoding: loading it as-is instead of re-encoding", + formConversionBytes, + ) + } val over = totalBytes - b.bytes out += Suggestion("a smaller model: weights must shrink by ≥ ${MemoryPlans.formatBytes(over)} (e.g. a lower-bit quantization of the same model)", over) return out @@ -218,13 +266,14 @@ public object MemoryPlans { * - heap headroom: [HEAP_HEADROOM_BYTES]. */ public fun plan(input: PlanInput, budget: Budget? = null): MemoryPlan { - val weights = input.weights.sumOf { it.bytes } + val weights = input.weights.sumOf { it.residentBytes } + val weightsAsStored = input.weights.sumOf { it.bytes } val g = input.geometry val kvElements = if (g != null) kvElements(g, input.ctx) else 0L val kv = input.kvMode.bytes(kvElements) val kvAlt = (if (input.kvMode == KvCacheMode.TURBOQUANT_4) KvCacheMode.BF16 else KvCacheMode.TURBOQUANT_4).bytes(kvElements) val forward = if (g != null) forwardBytes(g, input.ctx, input.prefillChunk) else 0L - return MemoryPlan(input, weights, kv, kvAlt, forward, HEAP_HEADROOM_BYTES, budget) + return MemoryPlan(input, weights, kv, kvAlt, forward, HEAP_HEADROOM_BYTES, budget, weightsAsStored) } public fun kvElements(g: ModelGeometry, ctx: Int): Long = diff --git a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/memory/plan/WeightFormResolver.kt b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/memory/plan/WeightFormResolver.kt index 4a7663235..0cbcb5ba9 100644 --- a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/memory/plan/WeightFormResolver.kt +++ b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/memory/plan/WeightFormResolver.kt @@ -72,3 +72,28 @@ public object WeightFormResolver { return "${tenths / 10}.${tenths % 10}" } } + +/** + * Every weight in this input resolved to the form it will take on a device described by [profile] + * (#1116). + * + * The point of resolving *before* planning: `MemoryPlans.plan` prices `PlanTensor.residentBytes`, + * so a resolved dequantization shows up in the table and in the fit check, instead of being + * discovered when the load runs out of memory. Weights whose format has no encoding — dense ones — + * resolve to a pass-through form and change nothing. + */ +@ExperimentalMemoryApi +public fun PlanInput.resolveWeightForms( + profile: PlannerProfile, + capabilities: KernelCapabilities, +): PlanInput = copy( + weights = weights.map { tensor -> + tensor.copy( + form = WeightFormResolver.resolve( + stored = tensor.format.encoding.takeUnless { it is sk.ainet.lang.tensor.storage.TensorEncoding.Dense }, + profile = profile, + capabilities = capabilities, + ), + ) + }, +) diff --git a/skainet-lang/skainet-lang-core/src/commonTest/kotlin/sk/ainet/lang/memory/plan/WeightFormPlanPricingTest.kt b/skainet-lang/skainet-lang-core/src/commonTest/kotlin/sk/ainet/lang/memory/plan/WeightFormPlanPricingTest.kt new file mode 100644 index 000000000..b825c789d --- /dev/null +++ b/skainet-lang/skainet-lang-core/src/commonTest/kotlin/sk/ainet/lang/memory/plan/WeightFormPlanPricingTest.kt @@ -0,0 +1,114 @@ +package sk.ainet.lang.memory.plan + +import sk.ainet.lang.memory.ExperimentalMemoryApi +import sk.ainet.lang.memory.Format +import sk.ainet.lang.tensor.storage.TensorEncoding +import sk.ainet.lang.types.FP32 +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertTrue + +/** + * #1116: a declared form changes the plan, so a dequantization is a line in a table rather than an + * OOM at load. + */ +@OptIn(ExperimentalMemoryApi::class) +class WeightFormPlanPricingTest { + + private val elements = 1L shl 20 // 1 Mi weights, so the numbers are legible + + private fun q4kTensor(form: WeightForm? = null): PlanTensor { + val format = Format(FP32, TensorEncoding.Q4_K) + return PlanTensor( + name = "blk.0.attn_q.weight", + id = null, + format = format, + elementCount = elements, + bytes = format.physicalBytes(elements)!!, + form = form, + ) + } + + private fun input(vararg weights: PlanTensor) = PlanInput( + modelName = "synthetic", architecture = "llama", weights = weights.toList(), + geometry = null, ctx = 2048, + ) + + @Test + fun `a weight held as stored costs what the file holds`() { + val stored = q4kTensor() + assertEquals(stored.bytes, stored.residentBytes, "no form resolved: the file's bytes are the cost") + + val kept = q4kTensor(WeightForm(EncodingRequest.KeepAsStored)) + assertEquals(kept.bytes, kept.residentBytes, "KeepAsStored: same") + } + + @Test + fun `a dequantized weight costs its dense size rather than its packed one`() { + val dequantized = q4kTensor(WeightForm(EncodingRequest.DequantizeTo(FP32))) + assertEquals(elements * 4, dequantized.residentBytes, "FP32 is four bytes an element") + assertTrue( + dequantized.residentBytes > dequantized.bytes * 7, + "Q4_K → FP32 is roughly 8×; got ${dequantized.bytes} → ${dequantized.residentBytes}", + ) + } + + @Test + fun `the plan totals the resolved size and reports what the conversion added`() { + val plan = MemoryPlans.plan(input(q4kTensor(WeightForm(EncodingRequest.DequantizeTo(FP32))))) + val stored = q4kTensor().bytes + + assertEquals(elements * 4, plan.weightsBytes, "the plan holds the dense weight") + assertEquals(stored, plan.weightsAsStoredBytes, "and remembers what the file held") + assertEquals(elements * 4 - stored, plan.formConversionBytes, "the difference is the conversion's price") + } + + @Test + fun `a plan that only fits as stored does not claim to fit once dequantized`() { + // The failure mode #1116 exists to prevent: budget checked against the file's size, then + // the load quadruples it. + val stored = q4kTensor() + val budget = Budget.of(stored.bytes + MemoryPlans.HEAP_HEADROOM_BYTES + 1) + + assertEquals(true, MemoryPlans.plan(input(stored), budget).fits, "as stored, it fits") + + val dequantized = MemoryPlans.plan(input(q4kTensor(WeightForm(EncodingRequest.DequantizeTo(FP32)))), budget) + assertEquals(false, dequantized.fits, "dequantized, it does not — and the plan says so before the load") + } + + @Test + fun `the suggestion names the conversion since nobody asked for it`() { + val budget = Budget.of(q4kTensor().bytes) + val plan = MemoryPlans.plan(input(q4kTensor(WeightForm(EncodingRequest.DequantizeTo(FP32)))), budget) + + val suggestion = plan.suggestions().firstOrNull { it.text.contains("kernels for the stored encoding") } + assertTrue(suggestion != null, "expected a suggestion about the conversion, got ${plan.suggestions()}") + assertEquals(plan.formConversionBytes, suggestion.savesBytes, "it saves exactly what it costs") + } + + @Test + fun `the rendered table shows the conversion and hides it when there is none`() { + val converted = MemoryPlans.plan(input(q4kTensor(WeightForm(EncodingRequest.DequantizeTo(FP32))))).render() + assertTrue(converted.contains("re-encoded at load"), converted) + + val asStored = MemoryPlans.plan(input(q4kTensor())).render() + assertTrue(!asStored.contains("re-encoded at load"), "an unconverted plan reads exactly as before:\n$asStored") + } + + @Test + fun `resolving an input prices what the target can actually feed`() { + val stored = input(q4kTensor()) + + val withKernels = stored.resolveWeightForms(PlannerProfile.DESKTOP, KernelCapabilities.EVERYTHING) + assertEquals( + MemoryPlans.plan(stored).weightsBytes, MemoryPlans.plan(withKernels).weightsBytes, + "a target that can feed Q4_K holds Q4_K, and the plan is unchanged", + ) + + val withoutKernels = stored.resolveWeightForms(PlannerProfile.DESKTOP, KernelCapabilities.DENSE_ONLY) + assertEquals( + elements * 4, MemoryPlans.plan(withoutKernels).weightsBytes, + "a target with only dense kernels holds FP32 — the same file, four times the memory", + ) + } +}