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
6 changes: 6 additions & 0 deletions skainet-backends/skainet-backend-api/build.gradle.kts
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,12 @@ kotlin {
}

sourceSets {
commonTest.dependencies {
// the module declares its targets by hand (no sk.ainet.multiplatform convention plugin),
// so kotlin-test is wired here
implementation(libs.kotlin.test)
}

commonMain.dependencies {
// Neutral backend API is an `api` re-export of the tensor op and
// storage interfaces already defined in skainet-lang-core. Any
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,115 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.Format
import sk.ainet.lang.memory.Scope
import sk.ainet.lang.memory.TensorView
import sk.ainet.lang.memory.trace.NoopTraceSink
import sk.ainet.lang.memory.trace.TraceEvent
import sk.ainet.lang.memory.trace.TraceSink
import sk.ainet.lang.memory.trace.kernel as traceKernel
import sk.ainet.lang.tensor.Shape

/**
* Kernel selection on declared descriptors instead of an `is`-ladder over Kotlin classes
* (SKEEP-003 §5.1). The order is: **normalize** the operands as zero-copy views (so rank-1 decode
* steps never reach a kernel written for rank 2 — the #993 root cause disappears), build the
* [KernelKey], look it up, insert **visible** adapters when a kernel cannot take an operand as it
* is, and fall back to the reference kernel, which is correct for every format because it decodes.
*
* Adapters allocate in the caller's [Scope] (a `Forward` scope in a generation loop) and are
* emitted as [TraceEvent.AdapterInserted] — the "hidden 12 GB" of #782 becomes a visible event.
*/
@ExperimentalMemoryApi
public object KernelDispatch {

private val kernels: MutableList<ViewKernel> = mutableListOf()

/** Register [kernel]; later registrations win for the same key (a pack can override the reference). */
public fun register(kernel: ViewKernel) {
kernels.removeAll { it.key == kernel.key && it.name == kernel.name }
kernels.add(0, kernel)
}

/** Every registered kernel, most recently registered first. */
public fun kernels(): List<ViewKernel> = kernels.toList()

/** The kernel registered for [key], or `null`. */
public fun find(key: KernelKey): ViewKernel? = kernels.firstOrNull { it.key == key }

public fun clearForTesting() { kernels.clear() }

/**
* Normalize a matmul operand pair to rank 2 as **views** (rule 5, §5.1 "rank handling happens
* once"): `[k]` becomes `[1, k]`, `[b, s, k]` becomes `[b*s, k]` when contiguous. Returns the
* normalized activation and the number of leading dims that were flattened, so the caller can
* reshape the result back.
*/
public fun normalizeActivation(a: TensorView): Pair<TensorView, IntArray> = when {
a.shape.rank == 1 -> a.unsqueeze(0) to intArrayOf()
a.shape.rank == 2 -> a to intArrayOf()
else -> {
val leading = IntArray(a.shape.rank - 1) { a.shape[it] }
require(a.isContiguous) { "flattening leading dims needs a contiguous activation; materialize first" }
var rows = 1
for (d in leading) rows *= d
a.reshapeContiguous(Shape(rows, a.shape[a.shape.rank - 1])) to leading
}
}

/**
* Select and run `matmul(a, b)`, writing into [out]. [scope] owns any adapter the selection
* needs; [sink] sees the kernel run and every adapter.
*
* @throws UnsupportedKernelException when neither a kernel nor the reference path can serve the key
*/
public fun matmul(
a: TensorView,
b: TensorView,
out: TensorView,
scope: Scope = Scope.Ambient,
sink: TraceSink = NoopTraceSink,
) {
val key = KernelKey.matmul(a, b)
val exact = find(key)
if (exact != null) {
runTraced(exact, listOf(a, b), out, sink)
return
}
// No exact kernel: adapt the operands a kernel would accept, then fall back to the reference,
// which reads any format through decoding get().
val adaptedA = adapt(a, scope, sink, "gather")
val reference = ReferenceMatmulKernel(KernelKey.matmul(adaptedA, b))
runTraced(reference, listOf(adaptedA, b), out, sink)
}

/** Materialize [view] into a dense contiguous view when it is strided; emits an adapter event. */
public fun adapt(view: TensorView, scope: Scope, sink: TraceSink, kind: String): TensorView {
if (view.isContiguous || view.layout.blocked) return view
val dense = view.materialize(Format.dense(view.format.dtype), scope)
if (sink.isEnabled) {
sink.emit(TraceEvent.AdapterInserted(kind, view.format, dense.format, dense.elementCount * view.format.dtype.sizeInBytes, view.id))
}
return dense
}

private fun runTraced(kernel: ViewKernel, inputs: List<TensorView>, out: TensorView, sink: TraceSink) {
if (!sink.isEnabled) { kernel.run(inputs, out); return }
sink.traceKernel(
op = kernel.key.op,
kernel = kernel.name,
inputs = inputs.map { it.id },
output = out.id,
bytesRead = inputs.sumOf { it.elementCount * it.format.dtype.sizeInBytes },
bytesWritten = out.elementCount * out.format.dtype.sizeInBytes,
) { kernel.run(inputs, out) }
}
}

/** A view of the same contiguous bytes under a different shape (rule 5: reshape is a view). */
@ExperimentalMemoryApi
public fun TensorView.reshapeContiguous(newShape: Shape): TensorView {
require(isContiguous) { "reshape needs a contiguous view" }
require(newShape.volume.toLong() == elementCount) { "reshape must keep the element count ($elementCount), got ${newShape.volume}" }
return TensorView(newShape, format, sk.ainet.lang.memory.Layout.rowMajor(newShape, format, layout.offsetElements), storage, id)
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.Format
import sk.ainet.lang.memory.TensorView
import sk.ainet.lang.tensor.storage.TensorEncoding

/**
* What a kernel consumes, declared rather than discovered (SKEEP-003 §0 *KernelKey*, §5.1): the op,
* the [Format] of each operand, how each operand is laid out, and the placement it needs. The
* dispatcher looks a key up instead of walking an `is`-ladder over `TensorData` subclasses — which
* is what made every quantisation bug a dispatch bug (#993, #991).
*
* Keys are values: equal keys select the same kernel, and a key prints as something a log or an
* `UnsupportedKernel` message can show: `matmul(F32/Dense(4B) contiguous × F32/Q4_K blocked) @host`.
*/
@ExperimentalMemoryApi
public data class KernelKey(
val op: String,
val operands: List<OperandKey>,
val placement: Placement = Placement.HOST,
) {
/** 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"

public companion object {
/** The key of `matmul(activation, weight)` as the two views describe themselves. */
public fun matmul(activation: TensorView, weight: TensorView, placement: Placement = Placement.HOST): KernelKey =
KernelKey("matmul", listOf(OperandKey.of(activation), OperandKey.of(weight)), placement)
}
}

/** One operand of a [KernelKey]: its [Format] plus the layout class the kernel must cope with. */
@ExperimentalMemoryApi
public data class OperandKey(val format: Format, val layout: LayoutClass) {
override fun toString(): String = "$format ${layout.name.lowercase()}"

public companion object {
/** Describe [view]: dense-and-gap-free is `CONTIGUOUS`, a packed layout is `BLOCKED`, anything else `STRIDED`. */
public fun of(view: TensorView): OperandKey {
val cls = when {
view.layout.blocked -> LayoutClass.BLOCKED
view.isContiguous -> LayoutClass.CONTIGUOUS
else -> LayoutClass.STRIDED
}
return OperandKey(view.format, cls)
}

/** A dense contiguous operand of [format] — the shape kernels prefer. */
public fun contiguous(format: Format): OperandKey = OperandKey(format, LayoutClass.CONTIGUOUS)
}
}

/**
* How an operand's bytes are arranged, as far as kernel selection cares: one gap-free run
* ([CONTIGUOUS]), a strided view over a larger buffer ([STRIDED]), or block-packed ([BLOCKED]).
* A kernel that declares `CONTIGUOUS` gets a gather adapter inserted for a `STRIDED` operand
* (§5.1) — the adapter is visible in the trace, never hidden inside a kernel.
*/
@ExperimentalMemoryApi
public enum class LayoutClass { CONTIGUOUS, STRIDED, BLOCKED }

/** Thrown when no registered kernel and no adapter chain can serve a key; lists what is registered. */
@ExperimentalMemoryApi
public class UnsupportedKernelException(
public val key: KernelKey,
public val candidates: List<String>,
message: String = "No kernel for $key" + if (candidates.isEmpty()) "" else "; registered: ${candidates.joinToString(", ")}",
) : IllegalArgumentException(message)

/** The encoding name a [KernelKey] uses for a format, matching `KernelProvider.supports`' dtype keys. */
@ExperimentalMemoryApi
public val Format.kernelEncodingName: String
get() = when (val e = encoding) {
is TensorEncoding.Dense -> dtype.name
else -> e.name
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.Format
import sk.ainet.lang.memory.TensorView

/**
* A registered kernel: a function `(inputs, out) -> Unit` behind a [KernelKey] (SKEEP-003 §0
* *Kernel*, §5.2). Custom kernels are written against views and registered — an author never
* touches a `TensorData` subclass.
*/
@ExperimentalMemoryApi
public interface ViewKernel {
/** What this kernel serves. */
public val key: KernelKey

/** A name for logs, traces and `UnsupportedKernel` messages (`scalar-reference`, `panama-q4k`, …). */
public val name: String

/** Run the kernel: [inputs] as described by [key], result written into [out]. */
public fun run(inputs: List<TensorView>, out: TensorView)
}

/**
* The reference matmul: correct for **any** pair of formats and layouts, because it reads through
* `TensorView.get()`, which decodes (rule 4). Slow by design — it is the fallback that makes an
* unsupported combination produce right numbers with a warning instead of a `ClassCastException`
* in layer 17 (#993). Registered for every key the dispatcher cannot serve better.
*
* `out = a × bᵀ` in the shapes SKaiNET's dispatch normalises to: `a` is `[m, k]`, `b` is `[n, k]`
* (a weight stored output-major, as GGUF does), `out` is `[m, n]`.
*/
@ExperimentalMemoryApi
public class ReferenceMatmulKernel(override val key: KernelKey) : ViewKernel {
override val name: String get() = "reference"

override fun run(inputs: List<TensorView>, out: TensorView) {
require(inputs.size == 2) { "matmul takes two operands, got ${inputs.size}" }
val a = inputs[0]; val b = inputs[1]
require(a.shape.rank == 2 && b.shape.rank == 2 && out.shape.rank == 2) { "reference matmul works on rank-2 views (normalise first)" }
val m = a.shape[0]; val k = a.shape[1]; val n = b.shape[0]
require(b.shape[1] == k) { "inner dimensions disagree: a is [${m}, ${k}], b is [${n}, ${b.shape[1]}]" }
require(out.shape[0] == m && out.shape[1] == n) { "out must be [$m, $n], was ${out.shape}" }
for (i in 0 until m) {
for (j in 0 until n) {
var acc = 0f
for (t in 0 until k) acc += a.get(i, t) * b.get(j, t)
out.set(i, j, value = acc)
}
}
}

public companion object {
/** The reference kernel for the formats of [a] and [b]. */
@ExperimentalMemoryApi
public fun forOperands(a: Format, b: Format): ReferenceMatmulKernel =
ReferenceMatmulKernel(KernelKey("matmul", listOf(OperandKey.contiguous(a), OperandKey.contiguous(b))))
}
}
Loading
Loading