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
10 changes: 10 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,16 @@

### Added

- **NEON body for the Q4_0 matmul kernel.** `skainet_q4_0_matmul` was the only priority
quant format without a SIMD path (scalar C only, while q8_0/q4k/q5k/q6k had NEON).
It now unpacks the split-layout nibbles with `vand`/`vshr`, re-centres in the signed
int8 domain, and widens to f32 FMA lanes — plain NEON with no dotprod/i8mm requirement,
so it runs on every AArch64 core, and the same block-outer/row-inner loop order as
q8_0 (sequential weight reads; per-row accumulation order unchanged). Verified: 27/27
kernel tests green under qemu-aarch64 (cross-built `-march=armv8.2-a+fp16+dotprod`,
K/N-bundled gcc 8.3), `fmla` confirmed in the archive's disassembly; a new
Kotlin/Native Q4_0 parity test closes the gap where the aarch64 lane had no Q4_0
coverage at all. Part of the mobile-kernels effort (#920).
- **`skainet-backend-jni-cpu`: Android JNI bridge for the native NEON kernels.** ART has no
`java.lang.foreign`, so the priority-100 FFM provider can never run on Android — until now
Android inference ran on the priority-0 scalar floor. The new AAR ships the same C kernel
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,8 +22,12 @@
* `/proc/cpuinfo` confirmed asimddp + fphp/asimdhp present and i8mm
* absent, matching the chosen -march (no +i8mm).
* The linked archive was confirmed to contain udot/sdot + fmla, i.e. the
* SIMD paths — not the scalar fallback — executed. bf16 and q4_0 have no
* SIMD paths — not the scalar fallback — executed. bf16 and fp16 have no
* NEON path (scalar only).
*
* q4_0 gained a plain-NEON body (nibble unpack + widen + vfmaq_f32, no
* dotprod requirement) later — parity-checked under qemu-aarch64 via the
* same linuxArm64Test lane (#920).
*/

#if defined(__ARM_NEON) || defined(__ARM_NEON__)
Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
#include "skainet_kernels.h"
#include "skainet_simd.h"

#include <stddef.h>
#include <stdint.h>
Expand All @@ -20,6 +21,20 @@
* Dequant per element: `(code - 8) * d`. The `- 8` bias centres the
* unsigned 4-bit code. Scale `d` is folded once after the block
* accumulator (cheaper than broadcasting it across every inner FMA).
*
* Loop order: block OUTER, output row INNER — see q8_0_matmul.c for the
* rationale. The weight is block-major, so for a fixed block consecutive
* `o` are 18 bytes apart: weight bytes are read sequentially instead of
* striding output_dim*18 per step. Per output row the blocks still
* accumulate in ascending order, so results are numerically identical to
* the previous row-outer form.
*
* NEON path (SKAINET_HAVE_NEON): unpack the 16 code bytes into two
* int8x16 nibble vectors (vand / vshr), re-centre with vsubq_s8(.., 8),
* widen to f32 lanes and FMA against the activation lanes — the same
* widen+vfmaq_f32 structure as q8_0_matmul.c, because activations are
* FP32 (an int8 dotprod would need int8 activations). Plain NEON only,
* no dotprod/i8mm requirement, so the body runs on every AArch64 core.
*/

/* Portable FP16 → FP32 conversion. Matches the Kotlin
Expand Down Expand Up @@ -68,27 +83,65 @@ SKAINET_API void skainet_q4_0_matmul(
const int32_t BLOCK_SIZE = 32;
const int32_t BYTES_PER_BLOCK = 18;
const int32_t blocks_per_input_dim = input_dim / BLOCK_SIZE;
float* SKAINET_RESTRICT out_base = output + output_offset;

for (int32_t o = 0; o < output_dim; ++o) out_base[o] = 0.0f;

for (int32_t block_idx = 0; block_idx < blocks_per_input_dim; ++block_idx) {
const float* SKAINET_RESTRICT input_block =
input + input_offset + (size_t) block_idx * BLOCK_SIZE;
const uint8_t* SKAINET_RESTRICT block =
weight + weight_byte_offset +
(size_t)(block_idx * output_dim) * BYTES_PER_BLOCK;

for (int32_t o = 0; o < output_dim; ++o) {
float acc = 0.0f;
for (int32_t block_idx = 0; block_idx < blocks_per_input_dim; ++block_idx) {
const uint8_t* SKAINET_RESTRICT block =
weight + weight_byte_offset +
(size_t)(block_idx * output_dim + o) * BYTES_PER_BLOCK;
for (int32_t o = 0; o < output_dim; ++o, block += BYTES_PER_BLOCK) {
uint16_t d_bits = (uint16_t) block[0] | ((uint16_t) block[1] << 8);
float d = skainet_q4_0_fp16_to_fp32(d_bits);
const uint8_t* SKAINET_RESTRICT codes = block + 2;
const float* SKAINET_RESTRICT input_block =
input + input_offset + (size_t) block_idx * BLOCK_SIZE;
float block_sum = 0.0f;
#ifdef SKAINET_HAVE_NEON
/* Split-nibble unpack: low nibbles are elements 0..15, high
* nibbles are elements 16..31. Re-centre by 8 in the signed
* int8 domain, then widen to f32 and FMA — same structure as
* the q8_0 NEON body. */
const uint8x16_t packed = vld1q_u8(codes);
const int8x16_t lo8 = vsubq_s8(
vreinterpretq_s8_u8(vandq_u8(packed, vdupq_n_u8(0x0F))),
vdupq_n_s8(8));
const int8x16_t hi8 = vsubq_s8(
vreinterpretq_s8_u8(vshrq_n_u8(packed, 4)),
vdupq_n_s8(8));
const int16x8_t lo16a = vmovl_s8(vget_low_s8(lo8)); /* elems 0..7 */
const int16x8_t lo16b = vmovl_s8(vget_high_s8(lo8)); /* elems 8..15 */
const int16x8_t hi16a = vmovl_s8(vget_low_s8(hi8)); /* elems 16..23 */
const int16x8_t hi16b = vmovl_s8(vget_high_s8(hi8)); /* elems 24..31 */
float32x4_t accv = vdupq_n_f32(0.0f);
accv = vfmaq_f32(accv, vld1q_f32(input_block + 0),
vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo16a))));
accv = vfmaq_f32(accv, vld1q_f32(input_block + 4),
vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo16a))));
accv = vfmaq_f32(accv, vld1q_f32(input_block + 8),
vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo16b))));
accv = vfmaq_f32(accv, vld1q_f32(input_block + 12),
vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo16b))));
accv = vfmaq_f32(accv, vld1q_f32(input_block + 16),
vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi16a))));
accv = vfmaq_f32(accv, vld1q_f32(input_block + 20),
vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi16a))));
accv = vfmaq_f32(accv, vld1q_f32(input_block + 24),
vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi16b))));
accv = vfmaq_f32(accv, vld1q_f32(input_block + 28),
vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi16b))));
block_sum = skainet_neon_hadd_f32(accv);
#else
for (int32_t k = 0; k < 16; ++k) {
int32_t lo = (int32_t)(codes[k] & 0x0F) - 8;
int32_t hi = (int32_t)(codes[k] >> 4) - 8;
block_sum += input_block[k] * (float) lo;
block_sum += input_block[k + 16] * (float) hi;
}
acc += block_sum * d;
#endif
out_base[o] += block_sum * d;
}
output[output_offset + o] = acc;
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
package sk.ainet.exec.kernel

import kotlin.math.abs
import kotlin.random.Random
import kotlin.test.Test
import kotlin.test.assertTrue

/**
* Proves the Kotlin/Native cinterop path: [NativeKnQ4_0MatmulKernel] (calling
* the C `skainet_q4_0_matmul` via cinterop, linked from libskainet_kernels.a)
* must agree with the commonMain [ScalarQ4_0MatmulKernel] reference within
* FMA + `-ffast-math` reassociation tolerance.
*
* Runs on linuxX64 (host archive: scalar/auto-vectorized) AND linuxArm64
* (cross-built archive: plain-NEON nibble-unpack body added in #920), so the
* aarch64 run checks the `SKAINET_HAVE_NEON` path in q4_0_matmul.c against
* the scalar reference. Q4_0 blocks are 32 elements / 18 bytes (FP16 `d` +
* 16 bytes of split-layout nibbles: low nibbles decode elements 0..15, high
* nibbles elements 16..31); full-range random code bytes exercise both
* nibble lanes and the `- 8` re-centring.
*/
class NativeKnQ4_0MatmulKernelParityTest {

private val blockSize = 32
private val bytesPerBlock = 18

private fun randomQ4_0Bytes(numBlocks: Int, seed: Int): ByteArray {
val rng = Random(seed)
val bytes = ByteArray(numBlocks * bytesPerBlock)
rng.nextBytes(bytes)
for (block in 0 until numBlocks) {
val base = block * bytesPerBlock
// 0x3C00 == 1.0f16 for the per-block scale so dequant stays finite.
bytes[base + 0] = 0x00.toByte()
bytes[base + 1] = 0x3C.toByte()
}
return bytes
}

private fun assertParity(inputDim: Int, outputDim: Int, seed: Int, tol: Float) {
val numBlocks = (inputDim / blockSize) * outputDim
val packed = randomQ4_0Bytes(numBlocks, seed)
val input = FloatArray(inputDim) { Random(seed + it).nextFloat() - 0.5f }

val refOut = FloatArray(outputDim)
ScalarQ4_0MatmulKernel.matmul(input, 0, packed, 0, inputDim, outputDim, refOut, 0)

val knOut = FloatArray(outputDim)
NativeKnQ4_0MatmulKernel.matmul(input, 0, packed, 0, inputDim, outputDim, knOut, 0)

for (o in 0 until outputDim) {
val diff = abs(refOut[o] - knOut[o])
val rel = diff / (abs(refOut[o]) + 1e-9f)
assertTrue(
diff <= tol || rel < 1e-4f,
"row $o diverged: scalar=${refOut[o]} cinterop=${knOut[o]} diff=$diff rel=$rel tol=$tol",
)
}
}

@Test
fun single_block_single_row() = assertParity(32, 1, 42, 1e-2f)

@Test
fun single_block_multi_row() = assertParity(32, 16, 7, 1e-2f)

@Test
fun multi_block_multi_row() = assertParity(1024, 64, 123, 2e-1f)

@Test
fun llm_typical_shape() = assertParity(4096, 64, 999, 2e0f)
}
Loading