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
Expand Up @@ -129,6 +129,16 @@ public final class sk/ainet/exec/kernel/PanamaVectorQ8_0MatmulKernel : sk/ainet/
public fun matmul ([FI[BIII[FI)V
}

public final class sk/ainet/exec/kernel/Q4_KQ8ActivationReferenceKernel : sk/ainet/backend/api/kernel/Q4KMatmulKernel {
public static final field INSTANCE Lsk/ainet/exec/kernel/Q4_KQ8ActivationReferenceKernel;
public fun matmul ([FI[BIII[FI)V
}

public final class sk/ainet/exec/kernel/Q6_KQ8ActivationReferenceKernel : sk/ainet/backend/api/kernel/Q6KMatmulKernel {
public static final field INSTANCE Lsk/ainet/exec/kernel/Q6_KQ8ActivationReferenceKernel;
public fun matmul ([FI[BIII[FI)V
}

public final class sk/ainet/exec/kernel/ScalarBf16MatmulKernel : sk/ainet/backend/api/kernel/Bf16MatmulKernel {
public static final field INSTANCE Lsk/ainet/exec/kernel/ScalarBf16MatmulKernel;
public fun getCodec ()Lsk/ainet/lang/types/NarrowFloatCodec;
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,141 @@
package sk.ainet.exec.kernel

import sk.ainet.backend.api.kernel.Q4KMatmulKernel
import kotlin.math.abs

/**
* Test-support reference for the native/FFM/Kotlin-Native/JNI Q4_K kernels'
* ggml-style int8 activation-quantization fast path (`skainet_q4k_matmul` in
* `q4k_matmul.c`) — deliberately NOT [ScalarQ4_KMatmulKernel]'s exact-float
* algorithm.
*
* [ScalarQ4_KMatmulKernel] and the native-family kernels compute genuinely
* different things (#944): the native side quantizes the input activation to
* int8 first (ggml's `block_q8_K`, `d_in = maxabs/127`, symmetric round +
* clamp to [-127,127]) before an integer dot product against the 4-bit
* weight codes, then applies the block's `d`/`dMin` scale/min — a small,
* deliberate, expected source of divergence against the exact-float scalar
* kernel that per-row or RMS-energy tolerances have to absorb. Comparing a
* native kernel's output against *this* kernel instead — same activation
* quantization, same integer dot, same scale/min application, transcribed
* byte-for-byte from the C — isolates genuine kernel bugs (wrong offsets,
* layout, scale decode, dispatch) from that expected loss: agreement should
* be tight (float-accumulation-order / rounding-tie noise only), not
* activation-quant noise. Use a tight tolerance against this kernel; keep
* using the RMS-energy gate against [ScalarQ4_KMatmulKernel] for the
* "is the intended lossy path still within its expected budget" check.
*
* Not registered with any [sk.ainet.backend.api.kernel.KernelRegistry] and
* not tuned for speed (allocates per call) — test support only, never a
* production dispatch target.
*/
public object Q4_KQ8ActivationReferenceKernel : Q4KMatmulKernel {

private const val BLOCK_SIZE = 256
private const val SUB_BLOCK = 32
private const val BYTES_PER_BLOCK = 144

override fun matmul(
input: FloatArray, inputOffset: Int,
weight: ByteArray, weightByteOffset: Int,
inputDim: Int, outputDim: Int,
output: FloatArray, outputOffset: Int,
) {
require(inputDim % BLOCK_SIZE == 0) {
"Q4_KQ8ActivationReferenceKernel: inputDim must be a multiple of $BLOCK_SIZE; got $inputDim"
}
if (outputDim == 0) return
if (inputDim == 0) { for (o in 0 until outputDim) output[outputOffset + o] = 0f; return }

val blocksPerInputDim = inputDim / BLOCK_SIZE
// Pre-quantize the whole input row to Q8 once, mirroring the C
// kernel's single quantization pass reused across every output row.
val q8 = IntArray(inputDim)
val dIn = FloatArray(blocksPerInputDim)
for (b in 0 until blocksPerInputDim) {
val base = inputOffset + b * BLOCK_SIZE
var maxAbs = 0f
for (i in 0 until BLOCK_SIZE) {
val a = abs(input[base + i])
if (a > maxAbs) maxAbs = a
}
if (maxAbs == 0f) {
dIn[b] = 0f
for (i in 0 until BLOCK_SIZE) q8[b * BLOCK_SIZE + i] = 0
continue
}
val inv = 127f / maxAbs
dIn[b] = maxAbs / 127f
for (i in 0 until BLOCK_SIZE) {
var v = roundHalfAwayFromZero(input[base + i] * inv)
if (v > 127) v = 127 else if (v < -127) v = -127
q8[b * BLOCK_SIZE + i] = v
}
}

val scaleIdx = IntArray(8)
val minIdx = IntArray(8)
for (o in 0 until outputDim) output[outputOffset + o] = 0f

for (blockIdx in 0 until blocksPerInputDim) {
val di = dIn[blockIdx]
val q8Base = blockIdx * BLOCK_SIZE

for (o in 0 until outputDim) {
val blockBase = weightByteOffset + (blockIdx * outputDim + o) * BYTES_PER_BLOCK
val d = decodeHalf(((weight[blockBase + 1].toInt() and 0xFF) shl 8) or (weight[blockBase].toInt() and 0xFF))
val dMin = decodeHalf(((weight[blockBase + 3].toInt() and 0xFF) shl 8) or (weight[blockBase + 2].toInt() and 0xFF))

// ggml get_scale_min_k4 over the 12 scale bytes — identical to
// ScalarQ4_KMatmulKernel and skainet_q4k_decode_scales.
val sc = blockBase + 4
for (sb in 0 until 4) {
scaleIdx[sb] = weight[sc + sb].toInt() and 0x3F
minIdx[sb] = weight[sc + sb + 4].toInt() and 0x3F
}
for (sb in 4 until 8) {
val low4S = weight[sc + sb + 4].toInt() and 0x0F
val high2S = (weight[sc + sb - 4].toInt() and 0xFF) ushr 6
scaleIdx[sb] = low4S or (high2S shl 4)
val low4M = (weight[sc + sb + 4].toInt() and 0xFF) ushr 4
val high2M = (weight[sc + sb].toInt() and 0xFF) ushr 6
minIdx[sb] = low4M or (high2M shl 4)
}

val codesOffset = blockBase + 16
var blockScaleDot = 0L
var blockMinSum = 0L
for (groupJ in 0 until 4) {
val qsRegion = codesOffset + groupJ * 32
val sbLo = 2 * groupJ
val sbHi = sbLo + 1
val q8LoBase = q8Base + sbLo * SUB_BLOCK
val q8HiBase = q8Base + sbHi * SUB_BLOCK
var dotLo = 0
var sumLo = 0
var dotHi = 0
var sumHi = 0
for (i in 0 until SUB_BLOCK) {
val b = weight[qsRegion + i].toInt() and 0xFF
val codeLo = b and 0x0F
val codeHi = b ushr 4
val aLo = q8[q8LoBase + i]
val aHi = q8[q8HiBase + i]
dotLo += aLo * codeLo
sumLo += aLo
dotHi += aHi * codeHi
sumHi += aHi
}
blockScaleDot += scaleIdx[sbLo].toLong() * dotLo + scaleIdx[sbHi].toLong() * dotHi
blockMinSum += minIdx[sbLo].toLong() * sumLo + minIdx[sbHi].toLong() * sumHi
}

output[outputOffset + o] += di * (d * blockScaleDot.toFloat() - dMin * blockMinSum.toFloat())
}
}
}

/** Round-half-away-from-zero, matching `lrintf`'s effect for the non-tie floats real inputs produce. */
private fun roundHalfAwayFromZero(x: Float): Int =
if (x >= 0f) (x + 0.5f).toInt() else -(-x + 0.5f).toInt()
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,138 @@
package sk.ainet.exec.kernel

import sk.ainet.backend.api.kernel.Q6KMatmulKernel

/**
* Test-support reference for the native/FFM/Kotlin-Native/JNI Q6_K kernels'
* ggml-style int8 activation-quantization fast path (`skainet_q6k_matmul` in
* `q6k_matmul.c`) — deliberately NOT [ScalarQ6_KMatmulKernel]'s exact-float
* algorithm.
*
* Same rationale as [Q4_KQ8ActivationReferenceKernel] (see its kdoc and
* #944): the native side quantizes the activation to int8 first
* (`d_in = maxabs/127`), unpacks the 6-bit weight code to a centered int8
* `code - 32`, and does 16 int8 dot-products per block (one per scale
* group) instead of an exact float dequant-then-dot. Comparing a native
* kernel against this transcription instead of the exact-float scalar
* reference isolates genuine kernel bugs from the expected, small
* activation-quantization loss — use a tight tolerance here, keep the
* RMS-energy gate against [ScalarQ6_KMatmulKernel].
*
* Not registered with any [sk.ainet.backend.api.kernel.KernelRegistry] and
* not tuned for speed — test support only.
*/
public object Q6_KQ8ActivationReferenceKernel : Q6KMatmulKernel {

private const val BLOCK_SIZE = 256
private const val BYTES_PER_BLOCK = 210
private const val QL_OFFSET = 0
private const val QH_OFFSET = 128
private const val SCALES_OFFSET = 192
private const val D_OFFSET = 208

override fun matmul(
input: FloatArray, inputOffset: Int,
weight: ByteArray, weightByteOffset: Int,
inputDim: Int, outputDim: Int,
output: FloatArray, outputOffset: Int,
) {
require(inputDim % BLOCK_SIZE == 0) {
"Q6_KQ8ActivationReferenceKernel: inputDim must be a multiple of $BLOCK_SIZE; got $inputDim"
}
if (outputDim == 0) return
if (inputDim == 0) { for (o in 0 until outputDim) output[outputOffset + o] = 0f; return }

val blocksPerInputDim = inputDim / BLOCK_SIZE
val q8 = IntArray(inputDim)
val dIn = FloatArray(blocksPerInputDim)
for (b in 0 until blocksPerInputDim) {
val base = inputOffset + b * BLOCK_SIZE
var maxAbs = 0f
for (i in 0 until BLOCK_SIZE) {
val a = kotlin.math.abs(input[base + i])
if (a > maxAbs) maxAbs = a
}
if (maxAbs == 0f) {
dIn[b] = 0f
for (i in 0 until BLOCK_SIZE) q8[b * BLOCK_SIZE + i] = 0
continue
}
val inv = 127f / maxAbs
dIn[b] = maxAbs / 127f
for (i in 0 until BLOCK_SIZE) {
var v = roundHalfAwayFromZero(input[base + i] * inv)
if (v > 127) v = 127 else if (v < -127) v = -127
q8[b * BLOCK_SIZE + i] = v
}
}

val codes = IntArray(BLOCK_SIZE)
for (o in 0 until outputDim) output[outputOffset + o] = 0f

for (blockIdx in 0 until blocksPerInputDim) {
val di = dIn[blockIdx]
val q8Base = blockIdx * BLOCK_SIZE

for (o in 0 until outputDim) {
val blockBase = weightByteOffset + (blockIdx * outputDim + o) * BYTES_PER_BLOCK
val d = decodeHalf(
((weight[blockBase + D_OFFSET + 1].toInt() and 0xFF) shl 8) or
(weight[blockBase + D_OFFSET].toInt() and 0xFF),
)

unpackCodes(weight, blockBase, codes)

// Σ_g sc[g] · Σ_{i∈g} q8[i]·codes[i], over 16 contiguous 16-element
// scale groups — matches skainet_q6k_weighted_dot_generic exactly
// (same group start/scale-index formula, same accumulation order).
var wdot = 0L
for (half in 0 until 2) {
for (k in 0 until 4) {
for (isb in 0 until 2) {
val start = half * 128 + 32 * k + isb * 16
val gs = half * 8 + isb + 2 * k
val sc = weight[blockBase + SCALES_OFFSET + gs].toInt() // signed int8
var dot = 0
for (j in 0 until 16) {
dot += q8[q8Base + start + j] * codes[start + j]
}
wdot += sc.toLong() * dot
}
}
}

output[outputOffset + o] += d * di * wdot.toFloat()
}
}
}

/**
* Unpack one 256-element Q6_K block into centered int8 codes (`code - 32`,
* range [-32,31]) in natural element order. Byte-for-byte transcription of
* `skainet_q6k_unpack_codes`.
*/
private fun unpackCodes(weight: ByteArray, blockBase: Int, codes: IntArray) {
val ql0 = blockBase + QL_OFFSET
val qh0 = blockBase + QH_OFFSET
for (half in 0 until 2) {
val ql = ql0 + half * 64
val qh = qh0 + half * 32
val out = half * 128
for (isb in 0 until 2) {
val lStart = isb * 16
for (l in lStart until lStart + 16) {
val qL0 = weight[ql + l].toInt() and 0xFF
val qL32 = weight[ql + l + 32].toInt() and 0xFF
val qH = weight[qh + l].toInt() and 0xFF
codes[out + l + 0] = ((qL0 and 0x0F) or ((qH and 0x03) shl 4)) - 32
codes[out + l + 32] = ((qL32 and 0x0F) or (((qH ushr 2) and 0x03) shl 4)) - 32
codes[out + l + 64] = ((qL0 ushr 4) or (((qH ushr 4) and 0x03) shl 4)) - 32
codes[out + l + 96] = ((qL32 ushr 4) or (((qH ushr 6) and 0x03) shl 4)) - 32
}
}
}
}

private fun roundHalfAwayFromZero(x: Float): Int =
if (x >= 0f) (x + 0.5f).toInt() else -(-x + 0.5f).toInt()
}
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@ import org.junit.Assert.assertNotNull
import org.junit.Assert.assertTrue
import org.junit.Test
import org.junit.runner.RunWith
import sk.ainet.exec.kernel.Q4_KQ8ActivationReferenceKernel
import sk.ainet.exec.kernel.Q6_KQ8ActivationReferenceKernel
import sk.ainet.exec.kernel.ScalarQ4_0MatmulKernel
import sk.ainet.exec.kernel.ScalarQ4_KMatmulKernel
import sk.ainet.exec.kernel.ScalarQ5_0MatmulKernel
Expand Down Expand Up @@ -221,6 +223,38 @@ class JniKernelParityTest {
reference = ScalarQ6_KMatmulKernel::matmul, jni = JniKernels::q6kMatmul,
)

/**
* Tight per-row parity against [Q4_KQ8ActivationReferenceKernel] /
* [Q6_KQ8ActivationReferenceKernel] (#944): these references perform the
* same int8 activation quantization the JNI kernel does, so agreement
* should be tight like the exact-float formats above — a wide
* divergence here indicates a real kernel/bridge bug, not the expected
* quantization loss the RMS-gated tests above absorb.
*/
@Test
fun q4k_parity_single_block_q8ActivationReference() = assertExactParity(
256, 16, 42, 1e-2f, 256, 144, conditionQ4K,
reference = Q4_KQ8ActivationReferenceKernel::matmul, jni = JniKernels::q4kMatmul,
)

@Test
fun q4k_parity_q8ActivationReference() = assertExactParity(
1024, 64, 123, 1e-2f, 256, 144, conditionQ4K,
reference = Q4_KQ8ActivationReferenceKernel::matmul, jni = JniKernels::q4kMatmul,
)

@Test
fun q6k_parity_single_block_q8ActivationReference() = assertExactParity(
256, 16, 42, 1e-2f, 256, 210, conditionQ6K,
reference = Q6_KQ8ActivationReferenceKernel::matmul, jni = JniKernels::q6kMatmul,
)

@Test
fun q6k_parity_q8ActivationReference() = assertExactParity(
1024, 64, 999, 1e-2f, 256, 210, conditionQ6K,
reference = Q6_KQ8ActivationReferenceKernel::matmul, jni = JniKernels::q6kMatmul,
)

@Test
fun smoke_roundtrip() {
val input = floatArrayOf(1f, 2f, 3f, 4f)
Expand Down
Loading
Loading