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,133 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.Format
import sk.ainet.lang.memory.I8Absmax
import sk.ainet.lang.memory.Storage
import sk.ainet.lang.memory.TensorView
import sk.ainet.lang.tensor.storage.TensorEncoding
import sk.ainet.lang.types.FP32

/**
* A hand-written `bitnet_gemv` supplied by a platform pack (SKEEP-003 §5.2, M2-F4).
*
* Deliberately array-shaped rather than view-shaped: the implementations are JNI or cinterop shims
* that pin primitive arrays, and keeping the SPI at that level means the pack modules carry no
* knowledge of `TensorView`. [TernaryKernelPacks] does the unwrapping once.
*
* The weight is canonical row-major `TQ2_0` — the order [sk.ainet.lang.memory.TernaryCodec] writes
* and a GGUF holds. That is *not* the block-major order the Q4_0…Q6_K SPI kernels take, which is
* why this one can be bridged into the view registry while those still wait on #973: there is no
* ambiguity here, and the parity test pins it.
*/
@ExperimentalMemoryApi
public interface BitNetGemvNative {
/** A name for logs and traces, e.g. `neon-dotprod`. */
public val name: String

/** `out[o] = activationScale * Σ_k code(k) · weight(o, k)` for one token. */
public fun gemvTq2_0(
activation: ByteArray,
activationOffset: Int,
activationScale: Float,
weight: ByteArray,
weightByteOffset: Int,
inputDim: Int,
outputDim: Int,
out: FloatArray,
outOffset: Int,
)
}

/**
* Installs the ternary kernels: the portable reference always, and a platform `bitnet_gemv` on top
* of it when one is available (M2-F4).
*
* Removing the native artifact is not an error and never a crash — the reference kernel is already
* registered, so dispatch keeps working at reference speed and the caller is *told* through [warn]
* rather than left to wonder why decode got slower.
*/
@ExperimentalMemoryApi
public object TernaryKernelPacks {

/** Capability a `bitnet_gemv` pack declares when it needs ARMv8.2 dot-product instructions. */
public const val CAPABILITY_DOTPROD: String = "dotprod"

private val ternaryEncodings = listOf(TensorEncoding.TQ1_0, TensorEncoding.TQ2_0, TensorEncoding.BITNET_B1_58)

/**
* @param native the platform kernel, or `null` when its artifact is absent
* @param capabilities what [native] needs; recorded in the key so a device without them never
* selects it (§5.2)
* @param warn where the "running without the native kernel" notice goes
* @return the name of the kernel that will serve TQ2_0 — the pack's, or the reference
*/
public fun install(
native: BitNetGemvNative? = null,
capabilities: Set<String> = emptySet(),
warn: (String) -> Unit = {},
): String {
for (encoding in ternaryEncodings) {
KernelDispatch.register(BitNetGemvKernel(BitNetGemvKernel.keyFor(Format(FP32, encoding))))
}
if (native == null) {
warn(
"bitnet_gemv: no native kernel available — using the portable reference. " +
"Add the NEON artifact for the fast path; nothing else changes.",
)
return BitNetGemvKernel(BitNetGemvKernel.keyFor(Format(FP32, TensorEncoding.TQ2_0))).name
}
val key = BitNetGemvKernel.keyFor(Format(FP32, TensorEncoding.TQ2_0)).copy(capabilities = capabilities)
val kernel = NativeBitNetGemvKernel(native, key)
KernelDispatch.register(kernel)
// Registering under the capability-free key too: the dispatcher builds its key from the
// operands, which say nothing about the CPU. The pack only installs itself on a device that
// *has* the capability, so the two keys select the same kernel there and neither exists on
// a device that does not.
KernelDispatch.register(NativeBitNetGemvKernel(native, BitNetGemvKernel.keyFor(Format(FP32, TensorEncoding.TQ2_0))))
return kernel.name
}
}

/**
* A [ViewKernel] over a [BitNetGemvNative]: unwraps the two views once per call and hands the
* native kernel the arrays it wants.
*
* Falls back to the reference for anything the native kernel does not take — a multi-row
* activation (prefill) or storage that is not a heap array — instead of failing: the fast path is
* an optimization, never a correctness requirement.
*/
@ExperimentalMemoryApi
public class NativeBitNetGemvKernel(
private val native: BitNetGemvNative,
override val key: KernelKey,
) : ViewKernel {

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

private val reference = BitNetGemvKernel(key)

override fun run(inputs: List<TensorView>, out: TensorView) {
val a = inputs[0]
val w = inputs[1]
val rows = a.shape[0]
val activationBytes = (a.storage as? Storage.Heap)?.bytes
val weightBytes = (w.storage as? Storage.Heap)?.bytes
val outFloats = (out.storage as? Storage.Heap)?.floats
if (rows != 1 || activationBytes == null || weightBytes == null || outFloats == null) {
reference.run(inputs, out)
return
}
native.gemvTq2_0(
activation = activationBytes,
activationOffset = (a.storage as Storage.Heap).arrayOffset,
activationScale = I8Absmax.scaleOf(a, 0),
weight = weightBytes,
weightByteOffset = (w.storage as Storage.Heap).arrayOffset,
inputDim = a.shape[1],
outputDim = w.shape[0],
out = outFloats,
outOffset = (out.storage as Storage.Heap).arrayOffset,
)
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.I8Absmax
import sk.ainet.lang.memory.Scope
import sk.ainet.lang.memory.Storage
import sk.ainet.lang.memory.TensorView
import sk.ainet.lang.memory.TernaryBlockDecoder
import sk.ainet.lang.memory.TernaryCodec
import sk.ainet.lang.tensor.Shape
import sk.ainet.lang.tensor.storage.TensorEncoding
import sk.ainet.lang.types.FP32
import kotlin.test.Test
import kotlin.test.assertTrue
import kotlin.time.TimeSource

/**
* The **fallback**'s cost, measured where the test runs (#1041, M2-A2).
*
* `bitnet_gemv/reference` is what a device gets when the NEON artifact is absent, so its time is
* the denominator of the "NEON is N× faster" claim. This prints milliseconds per call for a fixed
* shape and asserts only that it produced numbers — a speed assertion here would be a flake on
* shared CI hardware, and the acceptance measurement belongs on the reference device, where this
* same test is run from the Kotlin/Native binary.
*/
@OptIn(ExperimentalMemoryApi::class)
class BitNetGemvTimingTest {

private companion object {
const val K = 1024 // four TQ2_0 blocks per row
const val N = 256
const val CALLS = 5
}

@Test
fun referenceKernelThroughput() {
var seed = 7
val weightValues = FloatArray(N * K) {
seed = seed * 1103515245 + 12345
((seed ushr 16) % 3 - 1) * 0.5f
}
val bytes = TernaryCodec.encode(TensorEncoding.TQ2_0, weightValues)
val weight = TensorView.packed(
Storage.Heap.wrap(bytes), Shape(N, K), TensorEncoding.TQ2_0,
TernaryBlockDecoder(TensorEncoding.TQ2_0),
)
val activationFloats = FloatArray(K) {
seed = seed * 1103515245 + 12345
((seed ushr 16) % 2000 - 1000) / 1000f
}
val activation = I8Absmax.requantize(
TensorView.dense(Storage.Heap.wrap(activationFloats), Shape(1, K), FP32),
Scope.Ambient,
)
val out = TensorView.dense(Storage.Heap.floats(N), Shape(1, N), FP32)
val kernel = BitNetGemvKernel(BitNetGemvKernel.keyFor(weight.format))

kernel.run(listOf(activation, weight), out) // warm up
val mark = TimeSource.Monotonic.markNow()
repeat(CALLS) { kernel.run(listOf(activation, weight), out) }
val perCall = mark.elapsedNow().inWholeMicroseconds / CALLS.toDouble() / 1000.0

println("[bitnet_gemv] reference k=$K n=$N: $perCall ms/call (${N.toLong() * K} MACs)")
assertTrue(perCall >= 0.0)
var sawNonZero = false
for (o in 0 until N) if (out.get(0, o) != 0f) { sawNonZero = true; break }
assertTrue(sawNonZero, "the timed kernel must actually compute something")
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,162 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.Format
import sk.ainet.lang.memory.I8Absmax
import sk.ainet.lang.memory.Scope
import sk.ainet.lang.memory.Storage
import sk.ainet.lang.memory.TensorView
import sk.ainet.lang.memory.TernaryBlockDecoder
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.storage.TensorEncoding
import sk.ainet.lang.types.FP32
import kotlin.math.abs
import kotlin.test.AfterTest
import kotlin.test.BeforeTest
import kotlin.test.Test
import kotlin.test.assertEquals
import kotlin.test.assertTrue

/**
* #1041 (M2-F4): a platform `bitnet_gemv` takes over when its artifact is present, and its absence
* is a warning and a slower kernel — never a crash.
*
* The NEON kernel itself is C and is tested where it can run; what belongs here is the contract
* around it, which every target can check: who gets registered, what happens without the artifact,
* and that the native path is only taken for the shapes it declares.
*/
@OptIn(ExperimentalMemoryApi::class)
class TernaryKernelPacksTest {

private val k = 256
private val n = 4

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

/** A stand-in for the JNI kernel: records that it ran, and computes the same thing. */
private class FakeNative(override val name: String = "neon-dotprod") : BitNetGemvNative {
var calls: Int = 0
override fun gemvTq2_0(
activation: ByteArray, activationOffset: Int, activationScale: Float,
weight: ByteArray, weightByteOffset: Int,
inputDim: Int, outputDim: Int,
out: FloatArray, outOffset: Int,
) {
calls++
val codes = TernaryCodec.codes(TensorEncoding.TQ2_0, weight, outputDim * inputDim, weightByteOffset)
val blocks = inputDim / 256
for (o in 0 until outputDim) {
var acc = 0f
for (b in 0 until blocks) {
val scaleOffset = weightByteOffset + ((o * blocks + b) * 66) + 64
val d = sk.ainet.lang.types.Fp16Codec.decode(
(weight[scaleOffset].toInt() and 0xFF) or ((weight[scaleOffset + 1].toInt() and 0xFF) shl 8),
)
var partial = 0
for (i in 0 until 256) {
val code = codes[o * inputDim + b * 256 + i].toInt()
if (code != 0) {
val a = activation[activationOffset + b * 256 + i].toInt()
partial += if (code > 0) a else -a
}
}
acc += partial * d
}
out[outOffset + o] = acc * activationScale
}
}
}

private fun weight(): TensorView {
var seed = 5
val values = FloatArray(n * k) {
seed = seed * 1103515245 + 12345
((seed ushr 16) % 3 - 1) * 0.5f
}
val bytes = TernaryCodec.encode(TensorEncoding.TQ2_0, values)
return TensorView.packed(
Storage.Heap.wrap(bytes), Shape(n, k), TensorEncoding.TQ2_0,
TernaryBlockDecoder(TensorEncoding.TQ2_0),
)
}

private fun activation(rows: Int = 1): TensorView {
var seed = 9
val floats = FloatArray(rows * k) {
seed = seed * 1103515245 + 12345
((seed ushr 16) % 2000 - 1000) / 1000f
}
return I8Absmax.requantize(TensorView.dense(Storage.Heap.wrap(floats), Shape(rows, k), FP32), Scope.Ambient)
}

@Test
fun withoutTheArtifactTheReferenceServesAndTheCallerIsTold() {
val warnings = mutableListOf<String>()
val serving = TernaryKernelPacks.install(native = null, warn = { warnings += it })

assertEquals("bitnet_gemv/reference", serving)
assertEquals(1, warnings.size, "exactly one notice, not a crash: $warnings")
assertTrue(warnings.single().contains("portable reference"), warnings.single())

// and dispatch still works
val out = TensorView.dense(Storage.Heap.floats(n), Shape(1, n), FP32)
val sink = RecordingTraceSink()
KernelDispatch.matmul(activation(), weight(), out, Scope.Ambient, sink)
assertEquals("bitnet_gemv/reference", sink.eventsOf<TraceEvent.KernelRun>().single().kernel)
}

@Test
fun withTheArtifactTheNativeKernelTakesOverAndAgreesWithTheReference() {
val native = FakeNative()
val warnings = mutableListOf<String>()
val serving = TernaryKernelPacks.install(native, setOf(TernaryKernelPacks.CAPABILITY_DOTPROD)) { warnings += it }
assertEquals("bitnet_gemv/neon-dotprod", serving)
assertTrue(warnings.isEmpty(), "nothing to warn about: $warnings")

val w = weight()
val a = activation()
val fromNative = TensorView.dense(Storage.Heap.floats(n), Shape(1, n), FP32)
val sink = RecordingTraceSink()
KernelDispatch.matmul(a, w, fromNative, Scope.Ambient, sink)
assertEquals("bitnet_gemv/neon-dotprod", sink.eventsOf<TraceEvent.KernelRun>().single().kernel)
assertEquals(1, native.calls)

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

@Test
fun theNativePathIsSkippedForShapesItDoesNotTake() {
val native = FakeNative()
TernaryKernelPacks.install(native)
val w = weight()
// prefill: more than one row is not what the gemv takes — the reference runs instead
val out = TensorView.dense(Storage.Heap.floats(2 * n), Shape(2, n), FP32)
NativeBitNetGemvKernel(native, BitNetGemvKernel.keyFor(w.format)).run(listOf(activation(rows = 2), w), out)
assertEquals(0, native.calls, "a multi-row activation falls back rather than failing")
var nonZero = false
for (r in 0 until 2) for (o in 0 until n) if (out.get(r, o) != 0f) nonZero = true
assertTrue(nonZero, "and it still computed the answer")
}

@Test
fun theCapabilityIsRecordedInTheKey() {
val native = FakeNative()
TernaryKernelPacks.install(native, setOf(TernaryKernelPacks.CAPABILITY_DOTPROD))
val keys = KernelDispatch.kernels().filter { it.name.startsWith("bitnet_gemv/neon") }.map { it.key }
assertTrue(
keys.any { it.capabilities == setOf(TernaryKernelPacks.CAPABILITY_DOTPROD) },
"the pack declares what it needs: $keys",
)
assertTrue(keys.any { it.capabilities.isEmpty() }, "and is reachable from an operand-only key: $keys")
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ set(SKAINET_KERNEL_SOURCES
${SKAINET_KERNELS_ROOT}/src/q4k_matmul.c
${SKAINET_KERNELS_ROOT}/src/q5k_matmul.c
${SKAINET_KERNELS_ROOT}/src/q6k_matmul.c
${SKAINET_KERNELS_ROOT}/src/bitnet_gemv.c
)

set(SKAINET_JNI_SHIM ${CMAKE_CURRENT_SOURCE_DIR}/skainet_jni.c)
Expand Down
Loading
Loading