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
@@ -0,0 +1,52 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.plan.KernelCapabilities
import sk.ainet.lang.tensor.storage.TensorEncoding

/**
* [KernelCapabilities] answered by the kernels actually registered on this device (#1109).
*
* Deliberately thin. `KernelProvider.supports("matmul", [input, weight])` already exists and is
* already the contract providers override when they ship kernels beyond the built-in accessors —
* ternary packs among them — so asking it is asking the same source dispatch will ask. Inventing a
* second capability table beside that one is how the two drift apart.
*
* ## Both registries, because there are two
*
* A matmul kernel can reach this device by either of two routes, and a capability answer that knows
* about one of them is wrong about the other. [KernelRegistry] holds the provider SPI — the
* per-encoding accessors and `supports` — and is what the eager quantized paths consult.
* [KernelDispatch] holds `KernelKey`-addressed view kernels, and is where `KernelPacks.installReference`
* puts the reference FP32 GEMM that every target is supposed to have. Ask only the first and a
* target carrying nothing but the reference pack is reported as unable to multiply dense floats,
* which is both false and exactly the kind of drift this object exists to avoid.
*
* Availability matters as much as registration: a provider whose `isAvailable()` is false on this
* CPU cannot feed anything, whatever it declares. Dispatch kernels carry their requirements in
* their key's capability set instead, and are filtered when they are selected.
*/
@ExperimentalMemoryApi
public object RegistryKernelCapabilities : KernelCapabilities {

/** The activation dtype every packed matmul kernel in the tree takes. */
private const val FP32_KEY: String = "Float32"

override fun canFeedMatmul(encoding: TensorEncoding?): Boolean {
val weightKey = when (encoding) {
null -> FP32_KEY
is TensorEncoding.Dense -> FP32_KEY
else -> encoding.name
}
val fromProviders = KernelRegistry.providers().any { provider ->
provider.isAvailable() && provider.supports("matmul", listOf(FP32_KEY, weightKey))
}
return fromProviders || dispatchHasMatmulFor(weightKey)
}

/** Whether a registered [ViewKernel] computes a matmul whose weight operand is [weightKey]. */
private fun dispatchHasMatmulFor(weightKey: String): Boolean =
KernelDispatch.kernels().any { kernel ->
kernel.key.op == "matmul" && kernel.key.operands.any { it.format.kernelEncodingName == weightKey }
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.tensor.storage.TensorEncoding
import kotlin.test.AfterTest
import kotlin.test.Test
import kotlin.test.assertFalse
import kotlin.test.assertTrue

/**
* #1109 slice 1: the capability answer comes from the registry, so it cannot disagree with dispatch.
*/
@OptIn(ExperimentalMemoryApi::class)
class RegistryKernelCapabilitiesTest {

@AfterTest fun cleanup() { KernelDispatch.clearForTesting(); KernelRegistry.clearForTesting() }

/** A provider that declares one packed matmul and, optionally, is not available here. */
private class OneEncodingProvider(
private val weightKey: String,
private val available: Boolean = true,
) : KernelProvider {
override val name: String = "fake-$weightKey"
override val priority: Int = 100
override fun isAvailable(): Boolean = available
override fun matmulFp32(): Fp32MatmulKernel? = null
override fun supports(opName: String, dtypeKeys: List<String>): Boolean =
opName == "matmul" && dtypeKeys == listOf("Float32", weightKey)
}

@Test
fun anEncodingWithARegisteredKernelCanBeFed() {
KernelRegistry.clearForTesting()
KernelRegistry.register(OneEncodingProvider("Q4_K"))

assertTrue(RegistryKernelCapabilities.canFeedMatmul(TensorEncoding.Q4_K))
assertFalse(
RegistryKernelCapabilities.canFeedMatmul(TensorEncoding.Q6_K),
"nothing registered a Q6_K kernel, so the honest answer is no",
)
}

@Test
fun aProviderThatIsNotAvailableHereCannotFeedAnything() {
// Registration is a claim; availability is the fact. A pack compiled for a CPU feature this
// device lacks must not make the resolver keep an encoding nothing can compute (§5.2, #920).
KernelRegistry.clearForTesting()
KernelRegistry.register(OneEncodingProvider("Q4_K", available = false))

assertFalse(RegistryKernelCapabilities.canFeedMatmul(TensorEncoding.Q4_K))
}

@Test
fun aFeedableBlockedEncodingWantsKernelFeedOrder() {
KernelRegistry.clearForTesting()
KernelRegistry.register(OneEncodingProvider("Q8_0"))

assertTrue(
RegistryKernelCapabilities.wantsKernelFeedOrder(TensorEncoding.Q8_0),
"every packed matmul kernel in the tree reads input-block-major bytes (#973)",
)
assertFalse(
RegistryKernelCapabilities.wantsKernelFeedOrder(TensorEncoding.Q4_K),
"an encoding it cannot feed has no order preference to state",
)
}

@Test
fun denseIsAnsweredByTheFp32Kernel() {
KernelRegistry.clearForTesting()
assertFalse(RegistryKernelCapabilities.canFeedMatmul(null), "an empty registry can feed nothing")

KernelPacks.installReference()
assertTrue(RegistryKernelCapabilities.canFeedMatmul(null), "the reference pack always carries FP32")
}
}
91 changes: 91 additions & 0 deletions skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api
Original file line number Diff line number Diff line change
Expand Up @@ -1471,6 +1471,53 @@ public final class sk/ainet/lang/memory/plan/DeviceMemory {
public final class sk/ainet/lang/memory/plan/DeviceMemory$Companion {
}

public abstract interface class sk/ainet/lang/memory/plan/EncodingRequest {
}

public final class sk/ainet/lang/memory/plan/EncodingRequest$DequantizeTo : sk/ainet/lang/memory/plan/EncodingRequest {
public fun <init> (Lsk/ainet/lang/types/DType;)V
public final fun component1 ()Lsk/ainet/lang/types/DType;
public final fun copy (Lsk/ainet/lang/types/DType;)Lsk/ainet/lang/memory/plan/EncodingRequest$DequantizeTo;
public static synthetic fun copy$default (Lsk/ainet/lang/memory/plan/EncodingRequest$DequantizeTo;Lsk/ainet/lang/types/DType;ILjava/lang/Object;)Lsk/ainet/lang/memory/plan/EncodingRequest$DequantizeTo;
public fun equals (Ljava/lang/Object;)Z
public final fun getDtype ()Lsk/ainet/lang/types/DType;
public fun hashCode ()I
public fun toString ()Ljava/lang/String;
}

public final class sk/ainet/lang/memory/plan/EncodingRequest$KeepAsStored : sk/ainet/lang/memory/plan/EncodingRequest {
public static final field INSTANCE Lsk/ainet/lang/memory/plan/EncodingRequest$KeepAsStored;
public fun equals (Ljava/lang/Object;)Z
public fun hashCode ()I
public fun toString ()Ljava/lang/String;
}

public final class sk/ainet/lang/memory/plan/EncodingRequest$RequantizeTo : sk/ainet/lang/memory/plan/EncodingRequest {
public fun <init> (Lsk/ainet/lang/tensor/storage/TensorEncoding;)V
public final fun component1 ()Lsk/ainet/lang/tensor/storage/TensorEncoding;
public final fun copy (Lsk/ainet/lang/tensor/storage/TensorEncoding;)Lsk/ainet/lang/memory/plan/EncodingRequest$RequantizeTo;
public static synthetic fun copy$default (Lsk/ainet/lang/memory/plan/EncodingRequest$RequantizeTo;Lsk/ainet/lang/tensor/storage/TensorEncoding;ILjava/lang/Object;)Lsk/ainet/lang/memory/plan/EncodingRequest$RequantizeTo;
public fun equals (Ljava/lang/Object;)Z
public final fun getEncoding ()Lsk/ainet/lang/tensor/storage/TensorEncoding;
public fun hashCode ()I
public fun toString ()Ljava/lang/String;
}

public abstract interface class sk/ainet/lang/memory/plan/KernelCapabilities {
public static final field Companion Lsk/ainet/lang/memory/plan/KernelCapabilities$Companion;
public abstract fun canFeedMatmul (Lsk/ainet/lang/tensor/storage/TensorEncoding;)Z
public fun wantsKernelFeedOrder (Lsk/ainet/lang/tensor/storage/TensorEncoding;)Z
}

public final class sk/ainet/lang/memory/plan/KernelCapabilities$Companion {
public final fun getDENSE_ONLY ()Lsk/ainet/lang/memory/plan/KernelCapabilities;
public final fun getEVERYTHING ()Lsk/ainet/lang/memory/plan/KernelCapabilities;
}

public final class sk/ainet/lang/memory/plan/KernelCapabilities$DefaultImpls {
public static fun wantsKernelFeedOrder (Lsk/ainet/lang/memory/plan/KernelCapabilities;Lsk/ainet/lang/tensor/storage/TensorEncoding;)Z
}

public final class sk/ainet/lang/memory/plan/KvCacheMode : java/lang/Enum {
public static final field BF16 Lsk/ainet/lang/memory/plan/KvCacheMode;
public static final field FP32 Lsk/ainet/lang/memory/plan/KvCacheMode;
Expand Down Expand Up @@ -1758,6 +1805,50 @@ public final class sk/ainet/lang/memory/plan/Suggestion {
public fun toString ()Ljava/lang/String;
}

public final class sk/ainet/lang/memory/plan/WeightByteOrder : java/lang/Enum {
public static final field AS_STORED Lsk/ainet/lang/memory/plan/WeightByteOrder;
public static final field KERNEL_FEED Lsk/ainet/lang/memory/plan/WeightByteOrder;
public static fun getEntries ()Lkotlin/enums/EnumEntries;
public static fun valueOf (Ljava/lang/String;)Lsk/ainet/lang/memory/plan/WeightByteOrder;
public static fun values ()[Lsk/ainet/lang/memory/plan/WeightByteOrder;
}

public final class sk/ainet/lang/memory/plan/WeightForm {
public static final field Companion Lsk/ainet/lang/memory/plan/WeightForm$Companion;
public fun <init> ()V
public fun <init> (Lsk/ainet/lang/memory/plan/EncodingRequest;Lsk/ainet/lang/memory/plan/WeightByteOrder;Lsk/ainet/lang/memory/plan/WeightResidency;)V
public synthetic fun <init> (Lsk/ainet/lang/memory/plan/EncodingRequest;Lsk/ainet/lang/memory/plan/WeightByteOrder;Lsk/ainet/lang/memory/plan/WeightResidency;ILkotlin/jvm/internal/DefaultConstructorMarker;)V
public final fun component1 ()Lsk/ainet/lang/memory/plan/EncodingRequest;
public final fun component2 ()Lsk/ainet/lang/memory/plan/WeightByteOrder;
public final fun component3 ()Lsk/ainet/lang/memory/plan/WeightResidency;
public final fun copy (Lsk/ainet/lang/memory/plan/EncodingRequest;Lsk/ainet/lang/memory/plan/WeightByteOrder;Lsk/ainet/lang/memory/plan/WeightResidency;)Lsk/ainet/lang/memory/plan/WeightForm;
public static synthetic fun copy$default (Lsk/ainet/lang/memory/plan/WeightForm;Lsk/ainet/lang/memory/plan/EncodingRequest;Lsk/ainet/lang/memory/plan/WeightByteOrder;Lsk/ainet/lang/memory/plan/WeightResidency;ILjava/lang/Object;)Lsk/ainet/lang/memory/plan/WeightForm;
public fun equals (Ljava/lang/Object;)Z
public final fun getEncoding ()Lsk/ainet/lang/memory/plan/EncodingRequest;
public final fun getOrder ()Lsk/ainet/lang/memory/plan/WeightByteOrder;
public final fun getResidency ()Lsk/ainet/lang/memory/plan/WeightResidency;
public fun hashCode ()I
public final fun isPassThrough ()Z
public fun toString ()Ljava/lang/String;
}

public final class sk/ainet/lang/memory/plan/WeightForm$Companion {
public final fun getAS_STORED_ON_HEAP ()Lsk/ainet/lang/memory/plan/WeightForm;
}

public final class sk/ainet/lang/memory/plan/WeightFormResolver {
public static final field INSTANCE Lsk/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/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;
public static fun getEntries ()Lkotlin/enums/EnumEntries;
public static fun valueOf (Ljava/lang/String;)Lsk/ainet/lang/memory/plan/WeightResidency;
public static fun values ()[Lsk/ainet/lang/memory/plan/WeightResidency;
}

public final class sk/ainet/lang/memory/trace/CompositeTraceSink : sk/ainet/lang/memory/trace/TraceSink {
public fun <init> (Ljava/util/List;)V
public fun <init> ([Lsk/ainet/lang/memory/trace/TraceSink;)V
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
package sk.ainet.lang.memory.plan

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.blockSpec
import sk.ainet.lang.tensor.storage.TensorEncoding

/**
* What a backend's kernels can actually be fed (#1109).
*
* The question [WeightFormResolver] needs and nothing could answer before: *which encodings can
* this target compute a matmul on without help?* A backend that has no kernel for the format on
* disk will still produce correct output — by dequantizing, on every forward pass — so the absence
* is invisible until someone profiles it. Asking first turns that into a decision made once.
*
* Declared here rather than in `skainet-backend-api` because the resolver lives beside
* [PlannerProfile] and the two are used together, while the backend modules are downstream. The
* registry-backed implementation is over there, where the providers are.
*/
@ExperimentalMemoryApi
public interface KernelCapabilities {

/**
* Can this target feed a matmul on a weight encoded as [encoding], with FP32 activations?
*
* `null` means dense — the question is then whether an FP32 matmul kernel exists at all, which
* for any real backend it does.
*/
public fun canFeedMatmul(encoding: TensorEncoding?): Boolean

/**
* Does the kernel for [encoding] read input-block-major bytes?
*
* Every packed matmul kernel in the tree does (#973), so the default is "yes, if it is blocked
* and we can feed it". A backend whose packed kernels read canonical order overrides this and
* gets `AS_STORED` bytes instead of a pointless permutation.
*/
public fun wantsKernelFeedOrder(encoding: TensorEncoding): Boolean =
encoding.blockSpec != null && canFeedMatmul(encoding)

public companion object {

/** A target with dense kernels and nothing else — the conservative assumption. */
public val DENSE_ONLY: KernelCapabilities = object : KernelCapabilities {
override fun canFeedMatmul(encoding: TensorEncoding?): Boolean = encoding == null
}

/** A target that can feed anything. Useful in tests; true of no real backend. */
public val EVERYTHING: KernelCapabilities = object : KernelCapabilities {
override fun canFeedMatmul(encoding: TensorEncoding?): Boolean = true
}
}
}
Loading
Loading