diff --git a/skainet-backends/skainet-backend-jni-cpu/src/androidTest/kotlin/sk/ainet/exec/kernel/jni/JniKernelParityTest.kt b/skainet-backends/skainet-backend-jni-cpu/src/androidTest/kotlin/sk/ainet/exec/kernel/jni/JniKernelParityTest.kt index f75565fb4..bcc51ffad 100644 --- a/skainet-backends/skainet-backend-jni-cpu/src/androidTest/kotlin/sk/ainet/exec/kernel/jni/JniKernelParityTest.kt +++ b/skainet-backends/skainet-backend-jni-cpu/src/androidTest/kotlin/sk/ainet/exec/kernel/jni/JniKernelParityTest.kt @@ -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 { @@ -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 { 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 @@ -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 + } } diff --git a/skainet-backends/skainet-backend-jni-cpu/src/androidTest/kotlin/sk/ainet/exec/kernel/jni/SmolLm2DecodeBenchmark.kt b/skainet-backends/skainet-backend-jni-cpu/src/androidTest/kotlin/sk/ainet/exec/kernel/jni/SmolLm2DecodeBenchmark.kt new file mode 100644 index 000000000..604ce72a2 --- /dev/null +++ b/skainet-backends/skainet-backend-jni-cpu/src/androidTest/kotlin/sk/ainet/exec/kernel/jni/SmolLm2DecodeBenchmark.kt @@ -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 = 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) + } +} diff --git a/skainet-backends/skainet-backend-jni-cpu/src/main/kotlin/sk/ainet/exec/kernel/jni/JniKernels.kt b/skainet-backends/skainet-backend-jni-cpu/src/main/kotlin/sk/ainet/exec/kernel/jni/JniKernels.kt index 9575409be..e2ee855bd 100644 --- a/skainet-backends/skainet-backend-jni-cpu/src/main/kotlin/sk/ainet/exec/kernel/jni/JniKernels.kt +++ b/skainet-backends/skainet-backend-jni-cpu/src/main/kotlin/sk/ainet/exec/kernel/jni/JniKernels.kt @@ -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. */