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,93 @@
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 reference `ternary_f32_gemv` (#1138): **FP32 activations** against `BITNET_B1_58` ternary
* weights — exact, no activation quantization.
*
* This is the f32 sibling of [BitNetGemvKernel]. That kernel is the W1.58A8 path: it asks the
* dispatcher to requantize activations to int8-absmax first, trading ~1.5 % quantization error for
* `sdot`-friendly integer math. This one keeps the activations as they are — a ternary weight is
* a `+1`, `-1` or nothing, so the dot product is float adds and subtracts, and the result equals
* the FP32 matmul against the decoded weight bit-for-bit apart from summation order. It serves the
* **exact** dispatch key (FP32 dense × `BITNET_B1_58`), which `KernelDispatch.matmul` checks
* *before* the requantize branch — installing this pack short-circuits the int8 adapter with zero
* dispatcher changes.
*
* Only `BITNET_B1_58` is served: its payload (four codes per byte, low bit-pair first, in element
* order) is what the vendored NeoGPU LUT kernel (#1137) reads, and its scale is in-band. The GGML
* block types (TQ1_0/TQ2_0) have per-block scales and interleaved payloads — they keep their
* int8 path.
*
* Operands: `[rows, k]` FP32 dense activations × `[n, k]` `BITNET_B1_58` weights in canonical
* row-major order. Output `[rows, n]` FP32. The per-tensor scale is applied to the output —
* native implementations ([TernaryF32GemvNative]) do not see it.
*/
@ExperimentalMemoryApi
public class TernaryF32GemvKernel(override val key: KernelKey) : ViewKernel {

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

override fun run(inputs: List<TensorView>, out: TensorView) {
require(inputs.size == 2) { "ternary_f32_gemv takes (activation, weight), got ${inputs.size} operands" }
val a = inputs[0]
val w = inputs[1]
require(a.format.dtype == FP32) { "activation must be FP32, was ${a.format}" }
require(w.format.encoding == TensorEncoding.BITNET_B1_58) {
"weight must be ${TensorEncoding.BITNET_B1_58}, was ${w.format}"
}
require(a.shape.rank == 2 && w.shape.rank == 2 && out.shape.rank == 2) { "ternary_f32_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}" }
if (rows == 0 || n == 0) return

val bytes = weightBytes(w)
val byteOffset = (w.storage as Storage.Heap).arrayOffset
// Codes are hoisted out of the row loop, and the per-tensor scale out of everything.
val codes = TernaryCodec.codes(TensorEncoding.BITNET_B1_58, bytes, n * k, byteOffset)
val scale = TernaryCodec.bitNetScale(bytes, n * k, byteOffset)

for (r in 0 until rows) {
for (o in 0 until n) {
var acc = 0f
val base = o * k
for (i in 0 until k) {
when (codes[base + i].toInt()) {
1 -> acc += a.get(r, i)
-1 -> acc -= a.get(r, i)
2 -> acc += 2f * a.get(r, i) // byte code 3; loaders reject it, decode agrees
else -> Unit // zero weights cost nothing
}
}
out.set(r, o, value = acc * scale)
}
}
}

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

public companion object {
/** The exact key this kernel serves: FP32 dense contiguous × `BITNET_B1_58` row-major. */
public fun keyFor(): KernelKey = KernelKey(
op = "matmul",
operands = listOf(
OperandKey.contiguous(Format.dense(FP32)),
OperandKey(Format(FP32, TensorEncoding.BITNET_B1_58), LayoutClass.BLOCKED_ROW_MAJOR),
),
)
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,144 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.Storage
import sk.ainet.lang.memory.TensorView
import sk.ainet.lang.memory.TernaryCodec

/**
* The f32-activation ternary gemv a platform pack supplies (#1138) — the seam over the vendored
* NeoGPU LUT kernel (#1137, `skainet_ternary_f32_gemv`).
*
* Array-shaped like [BitNetGemvNative] and for the same reason: implementations are FFM, JNI or
* cinterop shims that pin primitive arrays; [TernaryF32KernelPack] does the view unwrapping once.
*
* The weight is the sequential `BITNET_B1_58` **payload** (four codes per byte, low bit-pair
* first, element order — byte-identical to what [TernaryCodec.encodeBitNet] writes). The
* per-tensor scale is NOT applied here — the wrapping view kernel owns it. `inputDim` must be a
* multiple of 4 (the packing is byte-per-4-elements per row).
*/
@ExperimentalMemoryApi
public interface TernaryF32GemvNative {
/** A name for logs and traces, e.g. `ffm`, `neon`. */
public val name: String

/** `out[o] = Σ_k activation(k) · code(o, k)` — one row, scale not applied. */
public fun gemvPacked(
activation: FloatArray,
activationOffset: Int,
weight: ByteArray,
weightByteOffset: Int,
inputDim: Int,
outputDim: Int,
out: FloatArray,
outOffset: Int,
)
}

/**
* Installs the exact FP32 × `BITNET_B1_58` path (#1138).
*
* Registration happens **only when a native kernel is present**, and that is deliberate: without
* the LUT kernel the portable way to multiply a ternary weight is the existing int8-requantize →
* `bitnet_gemv` path ([TernaryKernelPacks]), which stays untouched as the fallback. Registering a
* Kotlin f32 reference into dispatch would *shadow* that tuned path with a slower exact one; the
* f32 reference ([TernaryF32GemvKernel]) exists as the correctness oracle and the in-kernel
* fallback, not as a dispatch entry.
*
* Absence of the native artifact is a notice through [warn], never a crash — behavior is exactly
* today's.
*/
@ExperimentalMemoryApi
public object TernaryF32KernelPack {

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

/**
* @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 — the vendored kernel itself needs none beyond baseline NEON
* @param warn where the "running without the exact f32 path" notice goes
* @return the name of the kernel that will serve the exact key, or [NOT_INSTALLED]
*/
public fun install(
native: TernaryF32GemvNative? = null,
capabilities: Set<String> = emptySet(),
warn: (String) -> Unit = {},
): String {
if (native == null) {
warn(
"ternary_f32_gemv: no native kernel available — FP32×b1.58 matmuls keep the " +
"int8-requantize path. Add the native artifact for the exact f32 path; " +
"nothing else changes.",
)
return NOT_INSTALLED
}
val kernel = NativeTernaryF32ViewKernel(native, TernaryF32GemvKernel.keyFor().copy(capabilities = capabilities))
KernelDispatch.register(kernel)
// The dispatcher builds its key from the operands, which say nothing about the CPU — the
// capability-free key is the reachable one, the capability key documents the requirement.
// Same two-key pattern as TernaryKernelPacks.install.
KernelDispatch.register(NativeTernaryF32ViewKernel(native, TernaryF32GemvKernel.keyFor()))
return kernel.name
}
}

/**
* A [ViewKernel] over a [TernaryF32GemvNative]: unwraps the views once, loops the native gemv
* over the activation rows (prefill included — each row is an independent gemv), and applies the
* `BITNET_B1_58` per-tensor scale to what the native kernel wrote.
*
* Falls back to the reference for anything the native contract does not cover — non-heap storage,
* a strided view, or `k % 4 != 0` (the sequential packing crosses byte boundaries between rows
* then) — instead of failing: the fast path is an optimization, never a correctness requirement.
*/
@ExperimentalMemoryApi
public class NativeTernaryF32ViewKernel(
private val native: TernaryF32GemvNative,
override val key: KernelKey,
) : ViewKernel {

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

private val reference = TernaryF32GemvKernel(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
for (r in 0 until rows) {
native.gemvPacked(
activation = activationFloats,
activationOffset = aOffset + r * k,
weight = weightBytes,
weightByteOffset = wOffset,
inputDim = k,
outputDim = n,
out = outFloats,
outOffset = outOffset + r * n,
)
}
// The native kernel computes the unscaled codes-dot; the per-tensor scale lives in the
// weight's trailing FP32 and is applied once, here.
val scale = TernaryCodec.bitNetScale(weightBytes, n * k, wOffset)
if (scale != 1f) {
for (i in outOffset until outOffset + rows * n) outFloats[i] *= scale
}
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,100 @@
package sk.ainet.backend.api.kernel

import sk.ainet.lang.memory.ExperimentalMemoryApi
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.math.abs
import kotlin.test.Test
import kotlin.test.assertFailsWith
import kotlin.test.assertTrue

/**
* #1138: the f32 reference against the one thing it must equal — an FP32 matmul over the
* *decoded* weight ([TernaryCodec.decodeBitNet]). The codes-dot is exact; only summation
* order may differ.
*/
@OptIn(ExperimentalMemoryApi::class)
class TernaryF32GemvKernelTest {

private fun ternaryValues(count: Int, seed: Int): FloatArray {
var s = seed
return FloatArray(count) {
s = s * 1103515245 + 12345
((s ushr 16) % 3 - 1) * 0.7f
}
}

private fun weight(n: Int, k: Int, seed: Int = 5): TensorView {
val bytes = TernaryCodec.encodeBitNet(ternaryValues(n * k, seed))
return TensorView.packed(
Storage.Heap.wrap(bytes), Shape(n, k), TensorEncoding.BITNET_B1_58,
TernaryBlockDecoder(TensorEncoding.BITNET_B1_58, n * k),
)
}

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

@Test
fun referenceEqualsMatmulOverTheDecodedWeight() {
val rows = 2; val k = 96; val n = 5
val w = weight(n, k)
val a = activation(rows, k)
val out = TensorView.dense(Storage.Heap.floats(rows * n), Shape(rows, n), FP32)
TernaryF32GemvKernel(TernaryF32GemvKernel.keyFor()).run(listOf(a, w), out)

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

@Test
fun kIndivisibleByFourStillMatchesTheDecodedMatmul() {
// BITNET_B1_58 packs the flattened tensor, so k % 4 != 0 crosses byte
// boundaries between rows — the reference reads linear codes and must
// not care. (The native view kernel falls back to this path.)
val rows = 1; val k = 6; val n = 3
val w = weight(n, k, seed = 11)
val a = activation(rows, k, seed = 13)
val out = TensorView.dense(Storage.Heap.floats(rows * n), Shape(rows, n), FP32)
TernaryF32GemvKernel(TernaryF32GemvKernel.keyFor()).run(listOf(a, w), out)

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

@Test
fun mismatchedInnerDimensionsAreRejected() {
val w = weight(n = 2, k = 8)
val a = activation(rows = 1, k = 12)
val out = TensorView.dense(Storage.Heap.floats(2), Shape(1, 2), FP32)
assertFailsWith<IllegalArgumentException> {
TernaryF32GemvKernel(TernaryF32GemvKernel.keyFor()).run(listOf(a, w), out)
}
}
}
Loading
Loading