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,190 @@
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.TernaryCodec
import sk.ainet.lang.tensor.storage.TensorEncoding
import sk.ainet.lang.types.FP32

/**
* The fused multi-plane lm_head gemv a platform pack supplies (#1150) — the seam over the
* vendored NeoGPU `skainet_ternary_lmhead_stage1` (#1137). Array-shaped like the other ternary
* seams; the pack does the view unwrapping once.
*
* One call computes **planes [firstPlane, firstPlane+4)** of a `BITNET_PLANES` weight with the
* fused plane weights `{1, ⅓, ⅑, 1/27}` and the FP16 row scales applied:
*
* `out[o] = rowScale[o] · Σ_{q=0}^{3} (1/3^q) · Σ_k in[k] · code(firstPlane+q, o, k)`
*
* The full 8-plane matmul is two calls (`firstPlane` 0 and 4) combined as `s0 + s4 / 81` — which
* is how [NativeTernaryPlanesViewKernel] keeps the dispatch invariant *matmul == decoded matmul*.
* `inputDim % 4 == 0`; [rowScaleByteOffset] must be 2-byte aligned.
*/
@ExperimentalMemoryApi
public interface TernaryLmheadNative {
/** A name for logs and traces, e.g. `ffm`. */
public val name: String

public fun lmheadStage1(
activation: FloatArray,
activationOffset: Int,
weight: ByteArray,
planesByteOffset: Int,
planeStrideBytes: Int,
rowScaleByteOffset: Int,
inputDim: Int,
outputDim: Int,
out: FloatArray,
outOffset: Int,
)
}

/**
* Installs the exact FP32 × `BITNET_PLANES` path (#1150). Same contract as
* [TernaryF32KernelPack]: registration **only with a native kernel** — without it, dispatch falls
* back to the decoding reference matmul (correct, slow), told through [warn], never a crash.
*/
@ExperimentalMemoryApi
public object TernaryPlanesKernelPack {

/** What [install] returns when no native kernel is available and nothing was registered. */
public const val NOT_INSTALLED: String = "ternary_planes_matmul/not-installed"

public fun install(
native: TernaryLmheadNative? = null,
capabilities: Set<String> = emptySet(),
warn: (String) -> Unit = {},
): String {
if (native == null) {
warn(
"ternary_planes_matmul: no native kernel available — BITNET_PLANES matmuls decode " +
"through the reference path. Add the native artifact for the fused LUT path; " +
"nothing else changes.",
)
return NOT_INSTALLED
}
val kernel = NativeTernaryPlanesViewKernel(native, TernaryPlanesMatmulKernel.keyFor().copy(capabilities = capabilities))
KernelDispatch.register(kernel)
KernelDispatch.register(NativeTernaryPlanesViewKernel(native, TernaryPlanesMatmulKernel.keyFor()))
return kernel.name
}
}

/**
* The portable reference for FP32 × `BITNET_PLANES` — decodes each weight row (all 8 planes ×
* row scale) through [TernaryCodec] and accumulates in FP32. The correctness oracle and the
* in-kernel fallback; deliberately not a dispatch entry on its own.
*/
@ExperimentalMemoryApi
public class TernaryPlanesMatmulKernel(override val key: KernelKey) : ViewKernel {

override val name: String get() = "ternary_planes_matmul/reference"

override fun run(inputs: List<TensorView>, out: TensorView) {
require(inputs.size == 2) { "ternary_planes_matmul takes (activation, weight), got ${inputs.size}" }
val a = inputs[0]
val w = inputs[1]
require(w.format.encoding == TensorEncoding.BITNET_PLANES) {
"weight must be ${TensorEncoding.BITNET_PLANES}, was ${w.format}"
}
require(a.shape.rank == 2 && w.shape.rank == 2 && out.shape.rank == 2) { "ternary_planes_matmul is 2-D" }
val rows = a.shape[0]
val k = a.shape[1]
val n = w.shape[0]
require(w.shape[1] == k) { "inner dimensions differ: activation k=$k, weight k=${w.shape[1]}" }
require(out.shape[0] == rows && out.shape[1] == n) { "out must be [$rows, $n], was ${out.shape}" }
if (rows == 0 || n == 0) return

val heap = w.storage as? Storage.Heap
?: throw UnsupportedOperationException("ternary_planes_matmul reads weights from heap storage in this milestone")
val bytes = heap.bytes ?: throw UnsupportedOperationException("BITNET_PLANES weights need byte storage")
val byteOffset = heap.arrayOffset

val decodedRow = FloatArray(k)
for (o in 0 until n) {
TernaryCodec.decodeBitNetPlanesRow(bytes, n, k, o, decodedRow, 0, byteOffset)
for (r in 0 until rows) {
var acc = 0f
for (i in 0 until k) acc += a.get(r, i) * decodedRow[i]
out.set(r, o, value = acc)
}
}
}

public companion object {
/** The exact key: FP32 dense contiguous × `BITNET_PLANES` row-major. */
public fun keyFor(): KernelKey = KernelKey(
op = "matmul",
operands = listOf(
OperandKey.contiguous(Format.dense(FP32)),
OperandKey(Format(FP32, TensorEncoding.BITNET_PLANES), LayoutClass.BLOCKED_ROW_MAJOR),
),
)
}
}

/**
* A [ViewKernel] over a [TernaryLmheadNative]: two fused stage-1 calls per activation row —
* planes 0–3 and planes 4–7 — combined as `s0 + s4 / 81`, so the result is the **full 8-plane**
* matmul and the dispatch invariant *matmul == decoded matmul* holds exactly. (Stage-1-only
* scoring with top-k rescoring is an application-level decision made through the codec's row
* accessors, never through dispatch.)
*
* Falls back to the reference for non-heap storage, strided views, or `k % 4 != 0`.
*/
@ExperimentalMemoryApi
public class NativeTernaryPlanesViewKernel(
private val native: TernaryLmheadNative,
override val key: KernelKey,
) : ViewKernel {

override val name: String get() = "ternary_planes_matmul/${native.name}"

private val reference = TernaryPlanesMatmulKernel(key)

override fun run(inputs: List<TensorView>, out: TensorView) {
val a = inputs[0]
val w = inputs[1]
val rows = a.shape[0]
val k = a.shape[1]
val n = w.shape[0]
val activationFloats = (a.storage as? Storage.Heap)?.floats
val weightBytes = (w.storage as? Storage.Heap)?.bytes
val outFloats = (out.storage as? Storage.Heap)?.floats
if (k % 4 != 0 || activationFloats == null || weightBytes == null || outFloats == null ||
!a.isContiguous || !out.isContiguous
) {
reference.run(inputs, out)
return
}
if (rows == 0 || n == 0) return
val aOffset = (a.storage as Storage.Heap).arrayOffset
val wOffset = (w.storage as Storage.Heap).arrayOffset
val outOffset = (out.storage as Storage.Heap).arrayOffset
val planeStride = TensorEncoding.BITNET_PLANES.planeStrideBytes(n, k)
val scalesOffset = TensorEncoding.BITNET_PLANES.rowScalesByteOffset(n, k)
val high = FloatArray(n)
for (r in 0 until rows) {
// planes 0–3, row scales applied by the kernel
native.lmheadStage1(
activation = activationFloats, activationOffset = aOffset + r * k,
weight = weightBytes, planesByteOffset = wOffset,
planeStrideBytes = planeStride, rowScaleByteOffset = wOffset + scalesOffset,
inputDim = k, outputDim = n,
out = outFloats, outOffset = outOffset + r * n,
)
// planes 4–7, same fused weights — worth 1/3⁴ of the total
native.lmheadStage1(
activation = activationFloats, activationOffset = aOffset + r * k,
weight = weightBytes, planesByteOffset = wOffset + planeStride * 4,
planeStrideBytes = planeStride, rowScaleByteOffset = wOffset + scalesOffset,
inputDim = k, outputDim = n,
out = high, outOffset = 0,
)
val base = outOffset + r * n
for (o in 0 until n) outFloats[base + o] += high[o] / 81f
}
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,143 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.Scope
import sk.ainet.lang.memory.Storage
import sk.ainet.lang.memory.TensorView
import sk.ainet.lang.memory.TernaryCodec
import sk.ainet.lang.memory.trace.RecordingTraceSink
import sk.ainet.lang.memory.trace.TraceEvent
import sk.ainet.lang.tensor.Shape
import sk.ainet.lang.tensor.data.BitNetPlanesTensorData
import sk.ainet.lang.tensor.storage.TensorEncoding
import sk.ainet.lang.types.FP32
import kotlin.math.abs
import kotlin.random.Random
import kotlin.test.AfterTest
import kotlin.test.BeforeTest
import kotlin.test.Test
import kotlin.test.assertEquals
import kotlin.test.assertTrue

/**
* #1150: the `BITNET_PLANES` pack contract. The native seam computes 4 fused planes per call;
* the view kernel's two-call combination (`s0 + s4/81`) must equal the full 8-plane reference —
* the dispatch invariant *matmul == decoded matmul* holds exactly.
*/
@OptIn(ExperimentalMemoryApi::class)
class TernaryPlanesKernelPackTest {

private val k = 32
private val n = 5

@BeforeTest fun setUp() = KernelDispatch.clearForTesting()
@AfterTest fun tearDown() = KernelDispatch.clearForTesting()

/** Computes exactly the 4-plane fused contract, straight from the buffer layout. */
private class FakeNative(override val name: String = "fake") : TernaryLmheadNative {
var calls: Int = 0
override fun lmheadStage1(
activation: FloatArray, activationOffset: Int,
weight: ByteArray, planesByteOffset: Int,
planeStrideBytes: Int, rowScaleByteOffset: Int,
inputDim: Int, outputDim: Int,
out: FloatArray, outOffset: Int,
) {
calls++
val rowBytes = inputDim / 4
for (o in 0 until outputDim) {
val scaleBits = (weight[rowScaleByteOffset + o * 2].toInt() and 0xFF) or
((weight[rowScaleByteOffset + o * 2 + 1].toInt() and 0xFF) shl 8)
val scale = sk.ainet.lang.types.Fp16Codec.decode(scaleBits)
var acc = 0f
var w = 1f
for (q in 0 until 4) {
val base = planesByteOffset + q * planeStrideBytes + o * rowBytes
var dot = 0f
for (i in 0 until inputDim) {
val code = ((weight[base + i / 4].toInt() and 0xFF) shr ((i % 4) * 2)) and 3
dot += (code - 1) * activation[activationOffset + i]
}
acc += dot * w
w /= 3f
}
out[outOffset + o] = acc * scale
}
}
}

private fun weight(): TensorView {
val rng = Random(5)
val values = FloatArray(n * k) { (rng.nextFloat() - 0.5f) * 2f }
return BitNetPlanesTensorData.fromFloats(Shape(n, k), values).packedView
}

private fun activation(rows: Int = 1): TensorView {
val rng = Random(9)
return TensorView.dense(
Storage.Heap.wrap(FloatArray(rows * k) { rng.nextFloat() - 0.5f }),
Shape(rows, k), FP32,
)
}

@Test
fun withTheArtifactTheExactKeyServesAndEqualsTheFullPlaneReference() {
val native = FakeNative()
val serving = TernaryPlanesKernelPack.install(native)
assertEquals("ternary_planes_matmul/fake", serving)

val w = weight()
val a = activation()
val out = TensorView.dense(Storage.Heap.floats(n), Shape(1, n), FP32)
val sink = RecordingTraceSink()
KernelDispatch.matmul(a, w, out, Scope.Ambient, sink)
assertEquals("ternary_planes_matmul/fake", sink.eventsOf<TraceEvent.KernelRun>().single().kernel)
assertEquals(2, native.calls, "full result = planes 0–3 + planes 4–7 / 81 — two fused calls")

val fromReference = TensorView.dense(Storage.Heap.floats(n), Shape(1, n), FP32)
TernaryPlanesMatmulKernel(TernaryPlanesMatmulKernel.keyFor()).run(listOf(a, w), fromReference)
for (o in 0 until n) {
val got = out.get(0, o)
val want = fromReference.get(0, o)
assertTrue(abs(got - want) <= 1e-4f * maxOf(1f, abs(want)), "[$o]: $got vs $want")
}
}

@Test
fun withoutTheArtifactNothingIsRegisteredAndTheCallerIsTold() {
val warnings = mutableListOf<String>()
val serving = TernaryPlanesKernelPack.install(native = null, warn = { warnings += it })
assertEquals(TernaryPlanesKernelPack.NOT_INSTALLED, serving)
assertEquals(1, warnings.size)
assertTrue(KernelDispatch.kernels().isEmpty(), "nothing registered without the native kernel")
}

@Test
fun theReferenceEqualsTheDecodedMatmul() {
val w = weight()
val a = activation(rows = 2)
val out = TensorView.dense(Storage.Heap.floats(2 * n), Shape(2, n), FP32)
TernaryPlanesMatmulKernel(TernaryPlanesMatmulKernel.keyFor()).run(listOf(a, w), out)

val bytes = (w.storage as Storage.Heap).bytes!!
val decoded = TernaryCodec.decodeBitNetPlanes(bytes, n, k)
for (r in 0 until 2) for (o in 0 until n) {
var want = 0f
for (i in 0 until k) want += a.get(r, i) * decoded[o * k + i]
assertTrue(abs(out.get(r, o) - want) <= 1e-5f, "[$r,$o]: ${out.get(r, o)} vs $want")
}
}

@Test
fun encodingWithoutBlockSpecTriggersNoRequantizeAdapter() {
// BITNET_PLANES deliberately declares no activation hint: without the pack the dispatcher
// must fall to the decoding reference matmul, never the int8 requantize path.
val w = weight()
val a = activation()
val out = TensorView.dense(Storage.Heap.floats(n), Shape(1, n), FP32)
val sink = RecordingTraceSink()
KernelDispatch.matmul(a, w, out, Scope.Ambient, sink)
val kernel = sink.eventsOf<TraceEvent.KernelRun>().single().kernel
assertEquals("reference", kernel, "decoding reference serves without the pack")
}
}
Loading
Loading