diff --git a/skainet-backends/skainet-backend-api/src/commonMain/kotlin/sk/ainet/backend/api/kernel/KernelKey.kt b/skainet-backends/skainet-backend-api/src/commonMain/kotlin/sk/ainet/backend/api/kernel/KernelKey.kt index 3442aee2a..86da23fc2 100644 --- a/skainet-backends/skainet-backend-api/src/commonMain/kotlin/sk/ainet/backend/api/kernel/KernelKey.kt +++ b/skainet-backends/skainet-backend-api/src/commonMain/kotlin/sk/ainet/backend/api/kernel/KernelKey.kt @@ -19,12 +19,21 @@ public data class KernelKey( val op: String, val operands: List, val placement: Placement = Placement.HOST, + /** + * Platform capabilities a kernel requires (`vector`, `dotprod`, `i8mm`, `ffm`, …). Empty means + * "no special requirement" — the portable kernel. A pack registers its key *with* the + * capabilities it needs so a device that lacks them never selects it (§5.2, #920). + */ + val capabilities: Set = emptySet(), ) { /** Where the operands live — host memory today; a device backend adds its own (PRD non-goal for M1). */ public enum class Placement { HOST, DEVICE } - override fun toString(): String = - "$op(${operands.joinToString(" × ")})" + if (placement != Placement.HOST) " @${placement.name.lowercase()}" else " @host" + override fun toString(): String = buildString { + append(op); append('('); append(operands.joinToString(" × ")); append(')') + append(" @"); append(placement.name.lowercase()) + if (capabilities.isNotEmpty()) { append(" ["); append(capabilities.sorted().joinToString(",")); append(']') } + } public companion object { /** The key of `matmul(activation, weight)` as the two views describe themselves. */ diff --git a/skainet-backends/skainet-backend-api/src/commonMain/kotlin/sk/ainet/backend/api/kernel/KernelPacks.kt b/skainet-backends/skainet-backend-api/src/commonMain/kotlin/sk/ainet/backend/api/kernel/KernelPacks.kt new file mode 100644 index 000000000..173df2ce5 --- /dev/null +++ b/skainet-backends/skainet-backend-api/src/commonMain/kotlin/sk/ainet/backend/api/kernel/KernelPacks.kt @@ -0,0 +1,99 @@ +package sk.ainet.backend.api.kernel + +import sk.ainet.lang.memory.ExperimentalMemoryApi +import sk.ainet.lang.memory.Format +import sk.ainet.lang.memory.Storage +import sk.ainet.lang.memory.TensorView +import sk.ainet.lang.types.FP32 + +/** + * Wiring between the [KernelProvider] SPI (the platform packs: scalar, Panama/Vector API, native + * FFM, JNI/NEON) and the view-keyed [KernelDispatch] (SKEEP-003 §5.2). + * + * A pack keeps exposing its kernels through `KernelProvider`; this installs them as [ViewKernel]s + * under the keys the dispatcher looks up, so the generic path gets the fast kernel instead of the + * decoding reference whenever the operands' formats and layouts match what the pack declares. + * + * **Scope note (SKEEP-003 migration, #1029).** Only the dense FP32 kernel is bridged here. The + * packed (Q4_0…Q6_K) SPI kernels take their weight bytes in **block-major** order — the layout + * `DefaultCpuOpsBase.transposePackedBlocks` produces — while a packed `TensorView` describes the + * canonical row-major block order of the file. That contract is exactly what #973 reports as + * unwritten and contradictory across the engine and the converters, and getting it wrong is the + * silent-wrong-numbers class of #968/#971. Bridging the packed kernels therefore waits until #973 + * pins the byte order down; until then the packed fast paths stay on their existing (working) + * ladder in `DefaultCpuOps`/`DefaultCpuOpsJvm`, and the registry serves packed operands with the + * decoding reference kernel, which is correct for any layout. + */ +@ExperimentalMemoryApi +public object KernelPacks { + + /** Capability marker for a provider that needs an explicit vector unit (Panama, NEON, …). */ + public const val CAPABILITY_VECTOR: String = "vector" + + /** + * Install the kernels of [provider] (default: the best available one) into [KernelDispatch], + * plus the always-present reference kernel for dense FP32. Idempotent per provider name. + */ + public fun install(provider: KernelProvider? = KernelRegistry.bestAvailable()) { + installReference() + val p = provider ?: return + if (!p.isAvailable()) return + val fp32 = p.matmulFp32() ?: return + val dense = OperandKey.contiguous(Format.dense(FP32)) + // Two keys, one kernel: a weight normally reaches the dispatcher as a *transposed view* of a + // contiguous [k, n] buffer — strided by LayoutClass, but exactly what the SPI GEMM's stride + // arguments express. A contiguous weight is the [n, k]-in-memory case. + val strided = OperandKey(Format.dense(FP32), LayoutClass.STRIDED) + KernelDispatch.register(Fp32ViewMatmulKernel(p.name, fp32, KernelKey("matmul", listOf(dense, strided)))) + KernelDispatch.register(Fp32ViewMatmulKernel(p.name, fp32, KernelKey("matmul", listOf(dense, dense)))) + } + + /** The reference matmul for dense FP32 — always available, so a key is never unserved. */ + public fun installReference() { + val dense = OperandKey.contiguous(Format.dense(FP32)) + KernelDispatch.register(ReferenceMatmulKernel(KernelKey("matmul", listOf(dense, dense)))) + } +} + +/** + * A [ViewKernel] over an SPI [Fp32MatmulKernel]: both operands dense FP32 and contiguous, weight + * output-major (`[n, k]`, the shape SKaiNET's dispatch normalises to). Unwraps each view once — + * per the Phase-2 spike (#1016) — and calls the pack's strided GEMM. + */ +@ExperimentalMemoryApi +public class Fp32ViewMatmulKernel( + providerName: String, + private val kernel: Fp32MatmulKernel, + override val key: KernelKey, +) : ViewKernel { + override val name: String = "$providerName-fp32" + + override fun run(inputs: List, out: TensorView) { + require(inputs.size == 2) { "matmul takes two operands" } + val a = inputs[0]; val b = inputs[1] + val m = a.shape[0]; val k = a.shape[1]; val n = b.shape[0] + require(b.shape[1] == k) { "inner dimensions disagree: [${m}, ${k}] × [${n}, ${b.shape[1]}]" } + val aHeap = a.storage as? Storage.Heap ?: return fallback(inputs, out) + val bHeap = b.storage as? Storage.Heap ?: return fallback(inputs, out) + val oHeap = out.storage as? Storage.Heap ?: return fallback(inputs, out) + val aBuf = aHeap.floats ?: return fallback(inputs, out) + val bBuf = bHeap.floats ?: return fallback(inputs, out) + val oBuf = oHeap.floats ?: return fallback(inputs, out) + // The SPI GEMM reads the weight input-major: b[p][j] = bBuf[bOffset + p * bStride + j]. + // The dispatcher hands us the weight output-major ([n, k]) — which, when it is a transposed + // *view* of a contiguous [k, n] buffer, means strides[0] == 1 and strides[1] is that + // buffer's row stride. Anything else (a genuinely output-major buffer) would need a gather, + // so it goes to the reference kernel instead of being silently mis-indexed. + if (b.layout.strides[0] != 1) return fallback(inputs, out) + kernel.matmul( + a = aBuf, aOffset = aHeap.arrayOffset + a.layout.offsetElements.toInt(), aStride = a.layout.strides[0], + b = bBuf, bOffset = bHeap.arrayOffset + b.layout.offsetElements.toInt(), bStride = b.layout.strides[1], + out = oBuf, outOffset = oHeap.arrayOffset + out.layout.offsetElements.toInt(), outStride = out.layout.strides[0], + m = m, n = n, k = k, + ) + } + + private fun fallback(inputs: List, out: TensorView) { + ReferenceMatmulKernel(key).run(inputs, out) + } +} diff --git a/skainet-backends/skainet-backend-api/src/commonTest/kotlin/sk/ainet/backend/api/kernel/KernelPacksTest.kt b/skainet-backends/skainet-backend-api/src/commonTest/kotlin/sk/ainet/backend/api/kernel/KernelPacksTest.kt new file mode 100644 index 000000000..db84119b5 --- /dev/null +++ b/skainet-backends/skainet-backend-api/src/commonTest/kotlin/sk/ainet/backend/api/kernel/KernelPacksTest.kt @@ -0,0 +1,103 @@ +package sk.ainet.backend.api.kernel + +import sk.ainet.lang.memory.ExperimentalMemoryApi +import sk.ainet.lang.memory.Format +import sk.ainet.lang.memory.Storage +import sk.ainet.lang.memory.TensorView +import sk.ainet.lang.memory.trace.RecordingTraceSink +import sk.ainet.lang.memory.trace.TraceEvent +import sk.ainet.lang.tensor.Shape +import sk.ainet.lang.types.FP32 +import kotlin.math.abs +import kotlin.test.AfterTest +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertIs +import kotlin.test.assertTrue + +/** SKEEP-003 §5.2: platform packs register their kernels under view keys; the reference is always present. */ +@OptIn(ExperimentalMemoryApi::class) +class KernelPacksTest { + + @AfterTest fun cleanup() { KernelDispatch.clearForTesting(); KernelRegistry.clearForTesting() } + + /** A stand-in pack whose FP32 GEMM is a plain triple loop with the SPI's stride contract. */ + private class FakeProvider(override val name: String = "fake", override val priority: Int = 100) : KernelProvider { + var calls = 0 + override fun isAvailable(): Boolean = true + override fun matmulFp32(): Fp32MatmulKernel = object : Fp32MatmulKernel { + override fun matmul( + a: FloatArray, aOffset: Int, aStride: Int, + b: FloatArray, bOffset: Int, bStride: Int, + out: FloatArray, outOffset: Int, outStride: Int, + m: Int, n: Int, k: Int, + ) { + calls++ + for (i in 0 until m) for (j in 0 until n) { + var acc = 0f + for (p in 0 until k) acc += a[aOffset + i * aStride + p] * b[bOffset + p * bStride + j] + out[outOffset + i * outStride + j] = acc + } + } + } + } + + private fun view(shape: Shape, values: FloatArray): TensorView = + TensorView.dense(Storage.Heap.wrap(values), shape, FP32) + + @Test + fun theReferenceKernelIsAlwaysInstalled() { + KernelPacks.installReference() + val dense = OperandKey.contiguous(Format.dense(FP32)) + val k = KernelDispatch.find(KernelKey("matmul", listOf(dense, dense))) + assertEquals("reference", k?.name) + } + + @Test + fun aPackKernelServesTheDenseKeyAndAgreesWithTheReference() { + val provider = FakeProvider() + KernelRegistry.register(provider) + KernelPacks.install(provider) + + val a = view(Shape(2, 3), floatArrayOf(1f, 2f, 3f, 4f, 5f, 6f)) + // the weight as SKaiNET stores it, [k, n] = [3, 2], handed to the dispatcher transposed + val wBuf = floatArrayOf(1f, 0.5f, 2f, 1.5f, 3f, 2.5f) + val w = view(Shape(3, 2), wBuf).transpose() // [2, 3] output-major view + val out = view(Shape(2, 2), FloatArray(4)) + val sink = RecordingTraceSink() + KernelDispatch.matmul(a, w, out, sink = sink) + + assertTrue(provider.calls > 0, "the pack kernel must have run") + assertEquals("fake-fp32", assertIs(sink.events().single()).kernel) + // reference numbers + val expected = FloatArray(4) + for (i in 0 until 2) for (j in 0 until 2) { + var acc = 0f + for (p in 0 until 3) acc += a.get(i, p) * wBuf[p * 2 + j] + expected[i * 2 + j] = acc + } + for (i in expected.indices) assertTrue(abs(out.get(i / 2, i % 2) - expected[i]) < 1e-4f, "element $i: ${out.get(i / 2, i % 2)} vs ${expected[i]}") + } + + @Test + fun anOutputMajorWeightFallsBackToTheReferenceInsteadOfBeingMisIndexed() { + val provider = FakeProvider() + KernelRegistry.register(provider); KernelPacks.install(provider) + val a = view(Shape(1, 3), floatArrayOf(1f, 2f, 3f)) + // a genuinely output-major weight [n, k] (not a transposed view): strides [k, 1] + val w = view(Shape(2, 3), floatArrayOf(1f, 2f, 3f, 4f, 5f, 6f)) + val out = view(Shape(1, 2), FloatArray(2)) + KernelDispatch.matmul(a, w, out) + assertEquals(1f * 1 + 2f * 2 + 3f * 3, out.get(0, 0)) // reference semantics: out = a x wᵀ + assertEquals(1f * 4 + 2f * 5 + 3f * 6, out.get(0, 1)) + } + + @Test + fun keysCarryPlatformCapabilities() { + val dense = OperandKey.contiguous(Format.dense(FP32)) + val plain = KernelKey("matmul", listOf(dense, dense)) + val neon = KernelKey("matmul", listOf(dense, dense), capabilities = setOf("dotprod", KernelPacks.CAPABILITY_VECTOR)) + assertTrue(plain != neon, "capabilities are part of the key") + assertEquals("matmul(Float32/Dense(4B) contiguous × Float32/Dense(4B) contiguous) @host [dotprod,vector]", neon.toString()) + } +}