Skip to content

Commit 6891817

Browse files
Merge pull request #939 from SKaiNET-developers/feature/q4_0-neon-920
feat(backend): NEON body for the Q4_0 matmul kernel (#920)
2 parents f8d9f5a + c90c516 commit 6891817

4 files changed

Lines changed: 150 additions & 11 deletions

File tree

‎CHANGELOG.md‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,16 @@
44

55
### Added
66

7+
- **NEON body for the Q4_0 matmul kernel.** `skainet_q4_0_matmul` was the only priority
8+
quant format without a SIMD path (scalar C only, while q8_0/q4k/q5k/q6k had NEON).
9+
It now unpacks the split-layout nibbles with `vand`/`vshr`, re-centres in the signed
10+
int8 domain, and widens to f32 FMA lanes — plain NEON with no dotprod/i8mm requirement,
11+
so it runs on every AArch64 core, and the same block-outer/row-inner loop order as
12+
q8_0 (sequential weight reads; per-row accumulation order unchanged). Verified: 27/27
13+
kernel tests green under qemu-aarch64 (cross-built `-march=armv8.2-a+fp16+dotprod`,
14+
K/N-bundled gcc 8.3), `fmla` confirmed in the archive's disassembly; a new
15+
Kotlin/Native Q4_0 parity test closes the gap where the aarch64 lane had no Q4_0
16+
coverage at all. Part of the mobile-kernels effort (#920).
717
- **`skainet-backend-jni-cpu`: Android JNI bridge for the native NEON kernels.** ART has no
818
`java.lang.foreign`, so the priority-100 FFM provider can never run on Android — until now
919
Android inference ran on the priority-0 scalar floor. The new AAR ships the same C kernel

‎skainet-backends/skainet-backend-native-cpu/native/include/skainet_simd.h‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,8 +22,12 @@
2222
* `/proc/cpuinfo` confirmed asimddp + fphp/asimdhp present and i8mm
2323
* absent, matching the chosen -march (no +i8mm).
2424
* The linked archive was confirmed to contain udot/sdot + fmla, i.e. the
25-
* SIMD paths — not the scalar fallback — executed. bf16 and q4_0 have no
25+
* SIMD paths — not the scalar fallback — executed. bf16 and fp16 have no
2626
* NEON path (scalar only).
27+
*
28+
* q4_0 gained a plain-NEON body (nibble unpack + widen + vfmaq_f32, no
29+
* dotprod requirement) later — parity-checked under qemu-aarch64 via the
30+
* same linuxArm64Test lane (#920).
2731
*/
2832

2933
#if defined(__ARM_NEON) || defined(__ARM_NEON__)

‎skainet-backends/skainet-backend-native-cpu/native/src/q4_0_matmul.c‎

Lines changed: 63 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
#include "skainet_kernels.h"
2+
#include "skainet_simd.h"
23

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

2540
/* Portable FP16 → FP32 conversion. Matches the Kotlin
@@ -68,27 +83,65 @@ SKAINET_API void skainet_q4_0_matmul(
6883
const int32_t BLOCK_SIZE = 32;
6984
const int32_t BYTES_PER_BLOCK = 18;
7085
const int32_t blocks_per_input_dim = input_dim / BLOCK_SIZE;
86+
float* SKAINET_RESTRICT out_base = output + output_offset;
87+
88+
for (int32_t o = 0; o < output_dim; ++o) out_base[o] = 0.0f;
89+
90+
for (int32_t block_idx = 0; block_idx < blocks_per_input_dim; ++block_idx) {
91+
const float* SKAINET_RESTRICT input_block =
92+
input + input_offset + (size_t) block_idx * BLOCK_SIZE;
93+
const uint8_t* SKAINET_RESTRICT block =
94+
weight + weight_byte_offset +
95+
(size_t)(block_idx * output_dim) * BYTES_PER_BLOCK;
7196

72-
for (int32_t o = 0; o < output_dim; ++o) {
73-
float acc = 0.0f;
74-
for (int32_t block_idx = 0; block_idx < blocks_per_input_dim; ++block_idx) {
75-
const uint8_t* SKAINET_RESTRICT block =
76-
weight + weight_byte_offset +
77-
(size_t)(block_idx * output_dim + o) * BYTES_PER_BLOCK;
97+
for (int32_t o = 0; o < output_dim; ++o, block += BYTES_PER_BLOCK) {
7898
uint16_t d_bits = (uint16_t) block[0] | ((uint16_t) block[1] << 8);
7999
float d = skainet_q4_0_fp16_to_fp32(d_bits);
80100
const uint8_t* SKAINET_RESTRICT codes = block + 2;
81-
const float* SKAINET_RESTRICT input_block =
82-
input + input_offset + (size_t) block_idx * BLOCK_SIZE;
83101
float block_sum = 0.0f;
102+
#ifdef SKAINET_HAVE_NEON
103+
/* Split-nibble unpack: low nibbles are elements 0..15, high
104+
* nibbles are elements 16..31. Re-centre by 8 in the signed
105+
* int8 domain, then widen to f32 and FMA — same structure as
106+
* the q8_0 NEON body. */
107+
const uint8x16_t packed = vld1q_u8(codes);
108+
const int8x16_t lo8 = vsubq_s8(
109+
vreinterpretq_s8_u8(vandq_u8(packed, vdupq_n_u8(0x0F))),
110+
vdupq_n_s8(8));
111+
const int8x16_t hi8 = vsubq_s8(
112+
vreinterpretq_s8_u8(vshrq_n_u8(packed, 4)),
113+
vdupq_n_s8(8));
114+
const int16x8_t lo16a = vmovl_s8(vget_low_s8(lo8)); /* elems 0..7 */
115+
const int16x8_t lo16b = vmovl_s8(vget_high_s8(lo8)); /* elems 8..15 */
116+
const int16x8_t hi16a = vmovl_s8(vget_low_s8(hi8)); /* elems 16..23 */
117+
const int16x8_t hi16b = vmovl_s8(vget_high_s8(hi8)); /* elems 24..31 */
118+
float32x4_t accv = vdupq_n_f32(0.0f);
119+
accv = vfmaq_f32(accv, vld1q_f32(input_block + 0),
120+
vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo16a))));
121+
accv = vfmaq_f32(accv, vld1q_f32(input_block + 4),
122+
vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo16a))));
123+
accv = vfmaq_f32(accv, vld1q_f32(input_block + 8),
124+
vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo16b))));
125+
accv = vfmaq_f32(accv, vld1q_f32(input_block + 12),
126+
vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo16b))));
127+
accv = vfmaq_f32(accv, vld1q_f32(input_block + 16),
128+
vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi16a))));
129+
accv = vfmaq_f32(accv, vld1q_f32(input_block + 20),
130+
vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi16a))));
131+
accv = vfmaq_f32(accv, vld1q_f32(input_block + 24),
132+
vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi16b))));
133+
accv = vfmaq_f32(accv, vld1q_f32(input_block + 28),
134+
vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi16b))));
135+
block_sum = skainet_neon_hadd_f32(accv);
136+
#else
84137
for (int32_t k = 0; k < 16; ++k) {
85138
int32_t lo = (int32_t)(codes[k] & 0x0F) - 8;
86139
int32_t hi = (int32_t)(codes[k] >> 4) - 8;
87140
block_sum += input_block[k] * (float) lo;
88141
block_sum += input_block[k + 16] * (float) hi;
89142
}
90-
acc += block_sum * d;
143+
#endif
144+
out_base[o] += block_sum * d;
91145
}
92-
output[output_offset + o] = acc;
93146
}
94147
}
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,72 @@
1+
package sk.ainet.exec.kernel
2+
3+
import kotlin.math.abs
4+
import kotlin.random.Random
5+
import kotlin.test.Test
6+
import kotlin.test.assertTrue
7+
8+
/**
9+
* Proves the Kotlin/Native cinterop path: [NativeKnQ4_0MatmulKernel] (calling
10+
* the C `skainet_q4_0_matmul` via cinterop, linked from libskainet_kernels.a)
11+
* must agree with the commonMain [ScalarQ4_0MatmulKernel] reference within
12+
* FMA + `-ffast-math` reassociation tolerance.
13+
*
14+
* Runs on linuxX64 (host archive: scalar/auto-vectorized) AND linuxArm64
15+
* (cross-built archive: plain-NEON nibble-unpack body added in #920), so the
16+
* aarch64 run checks the `SKAINET_HAVE_NEON` path in q4_0_matmul.c against
17+
* the scalar reference. Q4_0 blocks are 32 elements / 18 bytes (FP16 `d` +
18+
* 16 bytes of split-layout nibbles: low nibbles decode elements 0..15, high
19+
* nibbles elements 16..31); full-range random code bytes exercise both
20+
* nibble lanes and the `- 8` re-centring.
21+
*/
22+
class NativeKnQ4_0MatmulKernelParityTest {
23+
24+
private val blockSize = 32
25+
private val bytesPerBlock = 18
26+
27+
private fun randomQ4_0Bytes(numBlocks: Int, seed: Int): ByteArray {
28+
val rng = Random(seed)
29+
val bytes = ByteArray(numBlocks * bytesPerBlock)
30+
rng.nextBytes(bytes)
31+
for (block in 0 until numBlocks) {
32+
val base = block * bytesPerBlock
33+
// 0x3C00 == 1.0f16 for the per-block scale so dequant stays finite.
34+
bytes[base + 0] = 0x00.toByte()
35+
bytes[base + 1] = 0x3C.toByte()
36+
}
37+
return bytes
38+
}
39+
40+
private fun assertParity(inputDim: Int, outputDim: Int, seed: Int, tol: Float) {
41+
val numBlocks = (inputDim / blockSize) * outputDim
42+
val packed = randomQ4_0Bytes(numBlocks, seed)
43+
val input = FloatArray(inputDim) { Random(seed + it).nextFloat() - 0.5f }
44+
45+
val refOut = FloatArray(outputDim)
46+
ScalarQ4_0MatmulKernel.matmul(input, 0, packed, 0, inputDim, outputDim, refOut, 0)
47+
48+
val knOut = FloatArray(outputDim)
49+
NativeKnQ4_0MatmulKernel.matmul(input, 0, packed, 0, inputDim, outputDim, knOut, 0)
50+
51+
for (o in 0 until outputDim) {
52+
val diff = abs(refOut[o] - knOut[o])
53+
val rel = diff / (abs(refOut[o]) + 1e-9f)
54+
assertTrue(
55+
diff <= tol || rel < 1e-4f,
56+
"row $o diverged: scalar=${refOut[o]} cinterop=${knOut[o]} diff=$diff rel=$rel tol=$tol",
57+
)
58+
}
59+
}
60+
61+
@Test
62+
fun single_block_single_row() = assertParity(32, 1, 42, 1e-2f)
63+
64+
@Test
65+
fun single_block_multi_row() = assertParity(32, 16, 7, 1e-2f)
66+
67+
@Test
68+
fun multi_block_multi_row() = assertParity(1024, 64, 123, 2e-1f)
69+
70+
@Test
71+
fun llm_typical_shape() = assertParity(4096, 64, 999, 2e0f)
72+
}

0 commit comments

Comments
 (0)