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,123 @@
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.memory.TernaryCodec
import sk.ainet.lang.memory.isTernary
import sk.ainet.lang.tensor.storage.TensorEncoding

/**
* The reference `bitnet_gemv` (SKEEP-003 §5.3, M2-F3): int8 activations against ternary weights,
* with **no multiplies in the inner loop**.
*
* A ternary weight is `-1`, `0` or `+1` times a scale, so the dot product is an add, a subtract or
* nothing at all — the whole point of 1.58-bit weights, and the shape a NEON kernel (#1041) will
* take. This one is plain Kotlin and runs everywhere: it is the thing the SIMD versions are
* checked against, so it is written for clarity, and the scales are factored out of the loop
* exactly as the vector kernels will factor them.
*
* Operands: `[rows, k]` activations in [I8Absmax.FORMAT] × `[n, k]` ternary weights in canonical
* (row-major block) order — the order [TernaryCodec] produces and GGUF stores. Output `[rows, n]`.
*
* The weight's codes are read once per call, not once per row: a decode step is one row against
* the whole matrix, so hoisting it is the difference between O(rows·n·k) decodes and O(n·k).
*/
@ExperimentalMemoryApi
public class BitNetGemvKernel(override val key: KernelKey) : ViewKernel {

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

override fun run(inputs: List<TensorView>, out: TensorView) {
require(inputs.size == 2) { "bitnet_gemv takes (activation, weight), got ${inputs.size} operands" }
val a = inputs[0]
val w = inputs[1]
require(a.format == I8Absmax.FORMAT) { "activation must be ${I8Absmax.FORMAT}, was ${a.format}" }
require(w.format.encoding.isTernary) { "weight must be ternary, was ${w.format}" }
require(a.shape.rank == 2 && w.shape.rank == 2 && out.shape.rank == 2) { "bitnet_gemv 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}" }

val encoding = w.format.encoding
val bytes = weightBytes(w)
val codes = TernaryCodec.codes(encoding, bytes, n * k)
val blockSize = blockSizeOf(encoding, n * k)

for (r in 0 until rows) {
val activation = I8Absmax.rowCodes(a, r)
val activationScale = I8Absmax.scaleOf(a, r)
for (o in 0 until n) {
var acc = 0f
var index = o * k
var offset = 0
// Walk the row block by block: within a block the weight scale is constant, so the
// inner loop only adds and subtracts activation codes (and skips the zeros).
while (offset < k) {
val run = minOf(blockSize - (index % blockSize), k - offset)
var partial = 0
for (i in 0 until run) {
when (codes[index + i].toInt()) {
1 -> partial += activation[offset + i].toInt()
-1 -> partial -= activation[offset + i].toInt()
else -> Unit // zero weights cost nothing
}
}
acc += partial * blockScale(encoding, bytes, (index + run - 1) / blockSize)
index += run
offset += run
}
out.set(r, o, value = acc * activationScale)
}
}
}

private fun weightBytes(w: TensorView): ByteArray {
val heap = w.storage as? Storage.Heap
?: throw UnsupportedOperationException("bitnet_gemv reads ternary weights from heap storage in this milestone")
return heap.bytes ?: throw UnsupportedOperationException("ternary weights need byte storage")
}

/** Elements per scale: a GGML block, or the whole tensor for the per-tensor BitNet encoding. */
private fun blockSizeOf(encoding: TensorEncoding, elements: Int): Int =
when (encoding) {
TensorEncoding.TQ1_0 -> TensorEncoding.TQ1_0.BLOCK_SIZE
TensorEncoding.TQ2_0 -> TensorEncoding.TQ2_0.BLOCK_SIZE
else -> elements
}

/** The scale of block [block] — per-block FP16 for the GGML types, one FP32 for BitNet. */
private fun blockScale(encoding: TensorEncoding, bytes: ByteArray, block: Int): Float = when (encoding) {
TensorEncoding.TQ1_0 -> fp16At(bytes, block * TensorEncoding.TQ1_0.BYTES_PER_BLOCK + 52)
TensorEncoding.TQ2_0 -> fp16At(bytes, block * TensorEncoding.TQ2_0.BYTES_PER_BLOCK + 64)
else -> TernaryCodec.bitNetScale(bytes, (bytes.size - TensorEncoding.BITNET_B1_58.SCALE_BYTES) * 4)
}

private fun fp16At(bytes: ByteArray, offset: Int): Float =
sk.ainet.lang.types.Fp16Codec.decode(
(bytes[offset].toInt() and 0xFF) or ((bytes[offset + 1].toInt() and 0xFF) shl 8),
)

public companion object {
/** The key this kernel serves for a ternary [weightFormat]. */
public fun keyFor(weightFormat: Format): KernelKey = KernelKey(
op = "matmul",
operands = listOf(
OperandKey.contiguous(I8Absmax.FORMAT),
OperandKey(weightFormat, LayoutClass.BLOCKED),
),
)

/** Register the reference kernel for every ternary encoding that carries its own bytes. */
public fun registerReference() {
for (encoding in listOf(TensorEncoding.TQ1_0, TensorEncoding.TQ2_0, TensorEncoding.BITNET_B1_58)) {
val format = Format(sk.ainet.lang.types.FP32, encoding)
KernelDispatch.register(BitNetGemvKernel(keyFor(format)))
}
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ 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.blockSpec
import sk.ainet.lang.memory.trace.NoopTraceSink
import sk.ainet.lang.memory.trace.TraceEvent
import sk.ainet.lang.memory.trace.TraceSink
Expand Down Expand Up @@ -76,13 +77,39 @@ public object KernelDispatch {
runTraced(exact, listOf(a, b), out, sink)
return
}
// The weight's encoding may *ask* for a different activation format — a ternary weight wants
// int8 with a per-token scale (`W1.58A8`, §5.3). Honour the request when a kernel exists for
// the requantized pair: the adapter costs bytes in the caller's scope every step, so it is
// allocated there and emitted as an AdapterInserted rather than hidden inside the kernel.
val wanted = b.format.encoding.blockSpec?.activation
if (wanted != null && wanted != a.format) {
val requantized = requantizeFor(wanted, a, scope, sink)
if (requantized != null) {
val ternaryKernel = find(KernelKey.matmul(requantized, b))
if (ternaryKernel != null) {
runTraced(ternaryKernel, listOf(requantized, 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)
}

/**
* Convert [activation] into the [wanted] activation format, or `null` when no adapter for it
* exists. Today the only one is the int8 absmax requantization the ternary kernels ask for.
*/
private fun requantizeFor(wanted: Format, activation: TensorView, scope: Scope, sink: TraceSink): TensorView? =
if (wanted == sk.ainet.lang.memory.I8Absmax.FORMAT && activation.shape.rank == 2) {
sk.ainet.lang.memory.I8Absmax.requantize(activation, scope, sink)
} else {
null
}

/** 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
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,213 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.Format
import sk.ainet.lang.memory.ForwardScope
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

/**
* #1040 (M2-F3): int8 activations against ternary weights, and the adapter that produces them.
*/
@OptIn(ExperimentalMemoryApi::class)
class BitNetGemvTest {

private val k = 256 // one TQ block
private val n = 4 // output rows

@BeforeTest fun setUp() {
KernelDispatch.clearForTesting()
BitNetGemvKernel.registerReference()
}

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

/** Deterministic ternary weight values, `[n, k]`, scaled so the block absmax is exact in FP16. */
private fun weightValues(seed: Int = 3): FloatArray {
var s = seed
return FloatArray(n * k) {
s = s * 1103515245 + 12345
((s ushr 16) % 3 - 1) * 0.5f
}
}

private fun activationValues(rows: Int, seed: Int = 11): FloatArray {
var s = seed
return FloatArray(rows * k) {
s = s * 1103515245 + 12345
((s ushr 16) % 2000 - 1000) / 1000f
}
}

private fun weightView(encoding: TensorEncoding, values: FloatArray): TensorView {
val bytes = TernaryCodec.encode(encoding, values)
val decoder = if (encoding == TensorEncoding.BITNET_B1_58) {
TernaryBlockDecoder(encoding, values.size) // one scale for the whole tensor
} else {
TernaryBlockDecoder(encoding)
}
return TensorView.packed(Storage.Heap.wrap(bytes), Shape(n, k), encoding, decoder)
}

/** `out[r, o] = Σ_k decoded_activation[r, k] * decoded_weight[o, k]` — the definition. */
private fun reference(activation: TensorView, weight: TensorView, rows: Int): FloatArray {
val out = FloatArray(rows * n)
for (r in 0 until rows) {
for (o in 0 until n) {
var acc = 0f
for (i in 0 until k) acc += I8Absmax.valueAt(activation, r, i) * weight.get(o, i)
out[r * n + o] = acc
}
}
return out
}

// --- the adapter ---------------------------------------------------------------------------

@Test
fun requantizationKeepsTheValuesAndPricesItself() {
val rows = 2
val values = activationValues(rows)
val dense = TensorView.dense(Storage.Heap.wrap(values), Shape(rows, k), FP32)
val sink = RecordingTraceSink()
val quantized = I8Absmax.requantize(dense, Scope.Ambient, sink)

assertEquals(I8Absmax.FORMAT, quantized.format)
for (r in 0 until rows) {
var amax = 0f
for (c in 0 until k) amax = maxOf(amax, abs(values[r * k + c]))
// Kotlin/JS computes Float arithmetic in double precision, so the same division can
// differ in the last bit between the test and the implementation: compare relatively.
val expectedScale = amax / 127f
assertTrue(
abs(expectedScale - I8Absmax.scaleOf(quantized, r)) <= 1e-6f * expectedScale,
"row $r scale should be absmax / 127 = $expectedScale, was ${I8Absmax.scaleOf(quantized, r)}",
)
val tolerance = I8Absmax.scaleOf(quantized, r)
for (c in 0 until k) {
assertTrue(
abs(values[r * k + c] - I8Absmax.valueAt(quantized, r, c)) <= tolerance,
"row $r col $c: ${values[r * k + c]} vs ${I8Absmax.valueAt(quantized, r, c)}",
)
}
}

val adapter = sink.eventsOf<TraceEvent.AdapterInserted>().single()
assertEquals("requantize-i8-absmax", adapter.kind)
assertEquals(I8Absmax.bytesFor(rows, k), adapter.bytes, "codes plus one scale per row")
assertEquals(rows.toLong() * k + rows * 4, adapter.bytes)
}

@Test
fun aZeroRowSurvivesQuantization() {
val dense = TensorView.dense(Storage.Heap.wrap(FloatArray(k)), Shape(1, k), FP32)
val quantized = I8Absmax.requantize(dense, Scope.Ambient)
assertEquals(0f, I8Absmax.scaleOf(quantized, 0), "no division by zero")
for (c in 0 until k) assertEquals(0f, I8Absmax.valueAt(quantized, row = 0, col = c))
}

@Test
fun theAdapterCostPerStepIsTheOneTheDesignPredicts() {
// §5.3 quotes ≈ 4 KB per decode step for a 2 B-parameter model (hidden 4096, one token).
val bytes = I8Absmax.bytesFor(rows = 1, cols = 4096)
assertEquals(4096L + 4, bytes)
assertTrue(bytes < 5 * 1024, "one token's activations must stay in the kilobytes: $bytes")
}

// --- the kernel ----------------------------------------------------------------------------

@Test
fun theReferenceKernelMatchesTheDefinitionForEveryTernaryEncoding() {
for (encoding in listOf(TensorEncoding.TQ1_0, TensorEncoding.TQ2_0, TensorEncoding.BITNET_B1_58)) {
val rows = 2
val weight = weightView(encoding, weightValues())
val activation = I8Absmax.requantize(
TensorView.dense(Storage.Heap.wrap(activationValues(rows)), Shape(rows, k), FP32),
Scope.Ambient,
)
val out = TensorView.dense(Storage.Heap.floats(rows * n), Shape(rows, n), FP32)
BitNetGemvKernel(BitNetGemvKernel.keyFor(weight.format)).run(listOf(activation, weight), out)

val expected = reference(activation, weight, rows)
for (r in 0 until rows) for (o in 0 until n) {
val got = out.get(r, o)
val want = expected[r * n + o]
assertTrue(
abs(got - want) <= 1e-3f * maxOf(1f, abs(want)),
"${encoding.name} [$r,$o]: $got vs $want",
)
}
}
}

@Test
fun zeroWeightsContributeNothing() {
// every weight zero → every output zero, whatever the activations are
val encoding = TensorEncoding.TQ2_0
val weight = weightView(encoding, FloatArray(n * k))
val activation = I8Absmax.requantize(
TensorView.dense(Storage.Heap.wrap(activationValues(1)), Shape(1, k), FP32),
Scope.Ambient,
)
val out = TensorView.dense(Storage.Heap.floats(n), Shape(1, n), FP32)
BitNetGemvKernel(BitNetGemvKernel.keyFor(weight.format)).run(listOf(activation, weight), out)
for (o in 0 until n) assertEquals(0f, out.get(0, o))
}

// --- dispatch ------------------------------------------------------------------------------

@Test
fun theDispatcherRequantizesTheActivationAndPicksTheTernaryKernel() {
val encoding = TensorEncoding.TQ2_0
val weight = weightView(encoding, weightValues())
val floats = activationValues(1)
val activation = TensorView.dense(Storage.Heap.wrap(floats), Shape(1, k), FP32)
val out = TensorView.dense(Storage.Heap.floats(n), Shape(1, n), FP32)

val sink = RecordingTraceSink()
val scope = ForwardScope(slabFloats = 4 * k, sink = sink, name = "decode")
KernelDispatch.matmul(activation, weight, out, scope, sink)

val kernels = sink.eventsOf<TraceEvent.KernelRun>()
assertEquals(1, kernels.size)
assertEquals("bitnet_gemv/reference", kernels.single().kernel, "a ternary weight selects the ternary kernel")
val adapter = sink.eventsOf<TraceEvent.AdapterInserted>().single()
assertEquals("requantize-i8-absmax", adapter.kind, "and the activation adapter is visible, not hidden")
assertEquals(I8Absmax.FORMAT, adapter.to)

// the numbers are the kernel's own
val quantized = I8Absmax.requantize(activation, Scope.Ambient)
val expected = reference(quantized, weight, rows = 1)
for (o in 0 until n) assertTrue(abs(out.get(0, o) - expected[o]) <= 1e-3f * maxOf(1f, abs(expected[o])), "[$o]")
scope.close()
}

@Test
fun aDenseWeightIsUntouchedByAnyOfThis() {
val weightFloats = FloatArray(n * k) { (it % 5) * 0.25f }
val weight = TensorView.dense(Storage.Heap.wrap(weightFloats), Shape(n, k), FP32)
val activation = TensorView.dense(Storage.Heap.wrap(activationValues(1)), Shape(1, k), FP32)
val out = TensorView.dense(Storage.Heap.floats(n), Shape(1, n), FP32)
val sink = RecordingTraceSink()
KernelDispatch.matmul(activation, weight, out, Scope.Ambient, sink)
assertTrue(
sink.eventsOf<TraceEvent.AdapterInserted>().none { it.kind == "requantize-i8-absmax" },
"only ternary formats ask for int8 activations",
)
}
}
Loading
Loading