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 @@ -11,15 +11,29 @@ import sk.ainet.exec.kernel.ScalarQ4_KMatmulKernel
import sk.ainet.exec.kernel.ScalarQ6_KMatmulKernel
import sk.ainet.exec.kernel.ScalarQ8_0MatmulKernel
import kotlin.math.abs
import kotlin.math.sqrt
import kotlin.random.Random

/**
* On-device parity: the JNI kernels must agree with the commonMain scalar
* references within FMA + `-ffast-math` reassociation tolerance. Runs on
* any arm64-v8a device (NEON baseline or v8.2+dotprod tier, whichever the
* loader picked) and on x86_64 emulators (scalar C paths) — the same
* parity contract the qemu-aarch64 lane enforces for the Kotlin/Native
* consumers of the identical C sources.
* references. Runs on any arm64-v8a device (NEON baseline or v8.2+dotprod
* tier, whichever the loader picked) and on x86_64 emulators (scalar C
* paths) — the same parity contract the `NativeKn*` qemu-aarch64 lane
* enforces for the Kotlin/Native consumers of the identical C sources.
*
* Two parity metrics, matching the kernel algorithm (see #944):
* - **q8_0 / q4_0**: exact-float kernels — the C dequantizes the weight and
* accumulates in FP32, so a per-row tolerance vs the scalar reference is
* right, catching any bit-level bug.
* - **q4_K / q6_K**: the C kernels quantize the *activation* to int8 first
* (ggml's `block_q8_K` fast path, faithful to `ggml_vec_dot_q4_K_q8_K`),
* which is deliberately lossy vs the exact-float scalar reference. Per-row
* relative error is meaningless on zero-mean random fixtures (a row whose
* true value is ~0 shows unbounded relative error from a tiny absolute
* one); the meaningful gate is the aggregate error ENERGY,
* RMS(error)/RMS(signal) — the same bar the `NativeKn*` K-format tests use.
* A structural bridge bug (wrong offset/layout/library) blows this by
* orders of magnitude; the intended quantization loss stays under it.
*/
@RunWith(AndroidJUnit4::class)
class JniKernelParityTest {
Expand All @@ -33,65 +47,162 @@ class JniKernelParityTest {
assertNotNull(JniKernels.variant)
}

private fun randomBlocks(numBlocks: Int, bytesPerBlock: Int, seed: Int): ByteArray {
/**
* The two-tier loader must pick the dotprod library on hardware that
* advertises it and the baseline elsewhere — the core of the #920
* runtime-dispatch design. Expectation is derived from the same
* `/proc/cpuinfo` signal the loader uses, so this passes on a dotprod
* arm64 device (→ V82_DOTPROD) and on an x86_64 emulator (→ BASELINE)
* without hard-coding either.
*/
@Test
fun loader_selects_tier_matching_cpu_features() {
val features = runCatching {
java.io.File("/proc/cpuinfo").useLines { lines ->
lines.firstOrNull { it.startsWith("Features") }
}
}.getOrNull().orEmpty()
val expectsDotprod = "asimddp" in features && ("asimdhp" in features || "fphp" in features)
val expected = if (expectsDotprod) {
JniKernels.Variant.V82_DOTPROD
} else {
JniKernels.Variant.BASELINE
}
assertEquals(
"loader tier must match CPU features (Features='$features')",
expected,
JniKernels.variant,
)
}

/**
* Random block bytes with each format's FP16 scale slots pinned to 1.0,
* mirroring the `NativeKn*` parity generators exactly (Q8_0/Q4_0 pin `d`
* at 0-1; Q4_K pins `d`+`dMin` at 0-3; Q6_K pins `d` at bytes 208-209).
*/
private fun randomBlocks(
numBlocks: Int, bytesPerBlock: Int, seed: Int,
condition: (bytes: ByteArray, base: Int) -> Unit,
): ByteArray {
val rng = Random(seed)
val bytes = ByteArray(numBlocks * bytesPerBlock)
rng.nextBytes(bytes)
for (block in 0 until numBlocks) {
val base = block * bytesPerBlock
bytes[base + 0] = 0x00.toByte() // FP16 scale = 1.0
bytes[base + 1] = 0x3C.toByte()
condition(bytes, block * bytesPerBlock)
}
return bytes
}

private fun assertParity(
inputDim: Int, outputDim: Int, seed: Int, tol: Float,
private fun fp16One(bytes: ByteArray, offset: Int) {
bytes[offset] = 0x00.toByte() // FP16 1.0 = 0x3C00, LE
bytes[offset + 1] = 0x3C.toByte()
}

private val conditionQ8_0: (ByteArray, Int) -> Unit = { b, base -> fp16One(b, base) }
private val conditionQ4_0: (ByteArray, Int) -> Unit = { b, base -> fp16One(b, base) }
private val conditionQ4K: (ByteArray, Int) -> Unit = { b, base -> fp16One(b, base); fp16One(b, base + 2) }
private val conditionQ6K: (ByteArray, Int) -> Unit = { b, base -> fp16One(b, base + 208) }

private fun run(
inputDim: Int, outputDim: Int, seed: Int,
blockSize: Int, bytesPerBlock: Int,
condition: (ByteArray, Int) -> Unit,
reference: (FloatArray, Int, ByteArray, Int, Int, Int, FloatArray, Int) -> Unit,
jni: (FloatArray, Int, ByteArray, Int, Int, Int, FloatArray, Int) -> Unit,
) {
): Pair<FloatArray, FloatArray> {
val numBlocks = (inputDim / blockSize) * outputDim
val packed = randomBlocks(numBlocks, bytesPerBlock, seed)
val packed = randomBlocks(numBlocks, bytesPerBlock, seed, condition)
val input = FloatArray(inputDim) { Random(seed + it).nextFloat() - 0.5f }

val refOut = FloatArray(outputDim)
reference(input, 0, packed, 0, inputDim, outputDim, refOut, 0)
val jniOut = FloatArray(outputDim)
jni(input, 0, packed, 0, inputDim, outputDim, jniOut, 0)
return refOut to jniOut
}

/** Exact-float parity (q8_0/q4_0): per-row absolute-or-relative tolerance. */
private fun assertExactParity(
inputDim: Int, outputDim: Int, seed: Int, tol: Float,
blockSize: Int, bytesPerBlock: Int,
condition: (ByteArray, Int) -> Unit,
reference: (FloatArray, Int, ByteArray, Int, Int, Int, FloatArray, Int) -> Unit,
jni: (FloatArray, Int, ByteArray, Int, Int, Int, FloatArray, Int) -> Unit,
) {
val (refOut, jniOut) = run(inputDim, outputDim, seed, blockSize, bytesPerBlock, condition, reference, jni)
for (o in 0 until outputDim) {
val diff = abs(refOut[o] - jniOut[o])
val rel = diff / (abs(refOut[o]) + 1e-9f)
assertTrue(
"row $o diverged: scalar=${refOut[o]} jni=${jniOut[o]} diff=$diff (variant=${JniKernels.variant})",
"row $o diverged: scalar=${refOut[o]} jni=${jniOut[o]} diff=$diff rel=$rel (variant=${JniKernels.variant})",
diff <= tol || rel < 1e-4f,
)
}
}

/**
* Aggregate RMS parity (q4_K/q6_K): the C kernel's int8 activation quant
* (#944) makes per-row float parity meaningless; bound the total error
* energy instead. Same gate as `NativeKnQ4KMatmulKernelParityTest`.
*/
private fun assertRmsParity(
inputDim: Int, outputDim: Int, seed: Int, tol: Float,
blockSize: Int, bytesPerBlock: Int,
condition: (ByteArray, Int) -> Unit,
reference: (FloatArray, Int, ByteArray, Int, Int, Int, FloatArray, Int) -> Unit,
jni: (FloatArray, Int, ByteArray, Int, Int, Int, FloatArray, Int) -> Unit,
) {
val (refOut, jniOut) = run(inputDim, outputDim, seed, blockSize, bytesPerBlock, condition, reference, jni)
var sqErr = 0.0
var sqSig = 0.0
for (o in 0 until outputDim) {
val d = (refOut[o] - jniOut[o]).toDouble()
sqErr += d * d
sqSig += refOut[o].toDouble() * refOut[o].toDouble()
}
val rmsErr = sqrt(sqErr / outputDim)
val rmsSig = sqrt(sqSig / outputDim)
val relRms = rmsErr / (rmsSig + 1e-9)
assertTrue(
"Q8-activation RMS parity exceeded: relRms=$relRms (rmsErr=$rmsErr rmsSig=$rmsSig) " +
"over $outputDim rows, tol=$AGG_REL_TOL (variant=${JniKernels.variant})",
relRms < AGG_REL_TOL || rmsErr < tol,
)
}

@Test
fun q80_parity() = assertParity(
1024, 64, 42, 2e-1f, 32, 34,
ScalarQ8_0MatmulKernel::matmul, JniKernels::q80Matmul,
fun q80_parity() = assertExactParity(
1024, 64, 42, 2e-1f, 32, 34, conditionQ8_0,
reference = ScalarQ8_0MatmulKernel::matmul, jni = JniKernels::q80Matmul,
)

@Test
fun q40_parity() = assertParity(
1024, 64, 7, 2e-1f, 32, 18,
ScalarQ4_0MatmulKernel::matmul, JniKernels::q40Matmul,
fun q40_parity() = assertExactParity(
1024, 64, 7, 2e-1f, 32, 18, conditionQ4_0,
reference = ScalarQ4_0MatmulKernel::matmul, jni = JniKernels::q40Matmul,
)

@Test
fun q4k_parity() = assertParity(
1024, 64, 123, 2e-1f, 256, 144,
ScalarQ4_KMatmulKernel::matmul, JniKernels::q4kMatmul,
fun q4k_parity_single_block() = assertRmsParity(
256, 16, 42, 1e-2f, 256, 144, conditionQ4K,
reference = ScalarQ4_KMatmulKernel::matmul, jni = JniKernels::q4kMatmul,
)

@Test
fun q6k_parity() = assertParity(
1024, 64, 999, 2e-1f, 256, 210,
ScalarQ6_KMatmulKernel::matmul, JniKernels::q6kMatmul,
fun q4k_parity() = assertRmsParity(
1024, 64, 123, 5e-2f, 256, 144, conditionQ4K,
reference = ScalarQ4_KMatmulKernel::matmul, jni = JniKernels::q4kMatmul,
)

@Test
fun q6k_parity_single_block() = assertRmsParity(
256, 16, 42, 1e-2f, 256, 210, conditionQ6K,
reference = ScalarQ6_KMatmulKernel::matmul, jni = JniKernels::q6kMatmul,
)

@Test
fun q6k_parity() = assertRmsParity(
1024, 64, 999, 5e-2f, 256, 210, conditionQ6K,
reference = ScalarQ6_KMatmulKernel::matmul, jni = JniKernels::q6kMatmul,
)

@Test
Expand All @@ -102,4 +213,10 @@ class JniKernelParityTest {
assertEquals(2f, output[0], 0f)
assertEquals(8f, output[3], 0f)
}

private companion object {
// Aggregate Q8-activation RMS-relative-error bound (uniform-random
// worst case) — same bar as the NativeKn* K-format parity tests.
const val AGG_REL_TOL = 0.03
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,129 @@
package sk.ainet.exec.kernel.jni

import android.util.Log
import androidx.test.ext.junit.runners.AndroidJUnit4
import org.junit.Assert.assertTrue
import org.junit.Test
import org.junit.runner.RunWith
import sk.ainet.exec.kernel.ScalarQ8_0MatmulKernel
import kotlin.random.Random

/**
* On-device decode-throughput projection for SmolLM2-135M-Instruct (Q8_0) —
* the acceptance measurement for #920 (field report: 1.0 tok/s scalar on
* Android; usable ≥ ~3 tok/s).
*
* Autoregressive decode of a 135M model is dominated by the per-layer
* projection mat-vecs plus the lm_head; attention/softmax/RoPE/sampling are
* negligible at this size and short context. This benchmark runs the ACTUAL
* JNI NEON kernels on the ACTUAL phone CPU at SmolLM2's real weight shapes,
* sums one token's matmul wall-clock across all 30 layers + lm_head, and
* reports projected tok/s for the JNI path and the scalar floor.
*
* It is a kernel-throughput projection, NOT an end-to-end generation (that
* needs the full transformers stack — transformers#272). But Q8_0 mat-vec is
* memory-bound, so timing at real shapes captures the real bottleneck.
*
* Results go to logcat under tag SKAINET_BENCH (dump after the run).
*/
@RunWith(AndroidJUnit4::class)
class SmolLm2DecodeBenchmark {

private val tag = "SKAINET_BENCH"

// SmolLM2-135M-Instruct config (HF): hidden 576, intermediate 1536,
// 30 layers, 9 heads / 3 kv heads, head_dim 64, vocab 49152.
private val hidden = 576
private val inter = 1536
private val layers = 30
private val kvDim = 3 * 64 // 192
private val vocab = 49152

/** (inputDim, outputDim, perTokenCount) for each decode mat-vec. */
private data class Shape(val inDim: Int, val outDim: Int, val perToken: Int, val name: String)

private fun shapes(): List<Shape> = listOf(
Shape(hidden, hidden, layers, "q_proj"),
Shape(hidden, kvDim, layers, "k_proj"),
Shape(hidden, kvDim, layers, "v_proj"),
Shape(hidden, hidden, layers, "o_proj"),
Shape(hidden, inter, layers, "gate_proj"),
Shape(hidden, inter, layers, "up_proj"),
Shape(inter, hidden, layers, "down_proj"),
Shape(hidden, vocab, 1, "lm_head"),
)

/** Q8_0 packed weight: (inDim/32)*outDim blocks × 34 bytes, scale pinned to 1.0. */
private fun q8Weight(inDim: Int, outDim: Int, seed: Int): ByteArray {
val blocks = (inDim / 32) * outDim
val bytes = ByteArray(blocks * 34)
Random(seed).nextBytes(bytes)
for (b in 0 until blocks) { bytes[b * 34] = 0x00; bytes[b * 34 + 1] = 0x3C }
return bytes
}

private inline fun timeMedianNs(iters: Int, warmup: Int, body: () -> Unit): Long {
repeat(warmup) { body() }
val samples = LongArray(iters)
for (i in 0 until iters) {
val t0 = System.nanoTime()
body()
samples[i] = System.nanoTime() - t0
}
samples.sort()
return samples[iters / 2]
}

@Test
fun projected_decode_tokens_per_second() {
assertTrue("JNI provider unavailable", JniKernelProvider.isAvailable())
val variant = JniKernels.variant
Log.i(tag, "=== SmolLM2-135M Q8_0 decode projection — device tier: $variant ===")

// Pre-warm both kernels so ART's JIT has compiled them before the
// first timed shape (otherwise the first scalar shape absorbs the
// C2 compile and reads ~5x slow).
run {
val w = q8Weight(576, 576, 1)
val i = FloatArray(576); val o = FloatArray(576)
repeat(50) { JniKernels.q80Matmul(i, 0, w, 0, 576, 576, o, 0) }
repeat(50) { ScalarQ8_0MatmulKernel.matmul(i, 0, w, 0, 576, 576, o, 0) }
}

var jniPerTokenNs = 0.0
var scalarPerTokenNs = 0.0

for (s in shapes()) {
val w = q8Weight(s.inDim, s.outDim, seed = s.name.hashCode())
val input = FloatArray(s.inDim) { Random(it).nextFloat() - 0.5f }
val out = FloatArray(s.outDim)

val jniNs = timeMedianNs(iters = 25, warmup = 5) {
JniKernels.q80Matmul(input, 0, w, 0, s.inDim, s.outDim, out, 0)
}
// Scalar is far slower; fewer iters to keep the run bounded.
val scalarNs = timeMedianNs(iters = 5, warmup = 2) {
ScalarQ8_0MatmulKernel.matmul(input, 0, w, 0, s.inDim, s.outDim, out, 0)
}

jniPerTokenNs += jniNs.toDouble() * s.perToken
scalarPerTokenNs += scalarNs.toDouble() * s.perToken
Log.i(
tag,
"%-10s %5d->%-5d ×%-2d jni=%7.1fµs scalar=%8.1fµs speedup=%.1fx".format(
s.name, s.inDim, s.outDim, s.perToken,
jniNs / 1000.0, scalarNs / 1000.0, scalarNs.toDouble() / jniNs,
),
)
}

val jniTokPerSec = 1e9 / jniPerTokenNs
val scalarTokPerSec = 1e9 / scalarPerTokenNs
Log.i(tag, "--------------------------------------------------------------")
Log.i(tag, "per-token matmul time: jni=%.2f ms scalar=%.2f ms".format(jniPerTokenNs / 1e6, scalarPerTokenNs / 1e6))
Log.i(tag, "PROJECTED DECODE: jni=%.2f tok/s scalar=%.2f tok/s (%.1fx)".format(jniTokPerSec, scalarTokPerSec, jniTokPerSec / scalarTokPerSec))
Log.i(tag, "usability gate (3 tok/s): jni %s".format(if (jniTokPerSec >= 3.0) "PASS ✅" else "FAIL ❌"))

assertTrue("projected tok/s must be positive", jniTokPerSec > 0)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -27,8 +27,16 @@ import java.io.File
*/
public object JniKernels {

/** Loaded library variant, or `null` when no variant could be loaded. */
public val variant: Variant? by lazy { loadVariant() }
/**
* Loaded library variant, or `null` when no variant could be loaded.
*
* Initialized EAGERLY in object init, not lazily: Kotlin object
* initialization runs on first access to ANY member, so a direct call
* to an `external` function is guaranteed to find the library loaded.
* (A lazy property would only load when `variant` itself is read —
* calling `smoke(...)` first would hit `UnsatisfiedLinkError`.)
*/
public val variant: Variant? = loadVariant()

public enum class Variant(public val libName: String) {
/** armv8-a baseline — every AArch64 device, plus x86_64 emulators. */
Expand Down
Loading