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

## [Unreleased]

### Added

- **Native Q5_0 / Q5_1 packed matmul kernels (FFM, Kotlin/Native, JNI).** 0.39.0 shipped
packed GGUF *loading* for Q5_0/Q5_1 plus scalar + Panama kernels, but the native tier had
no Q5_x kernels — on the JVM the registry cascaded to Panama (50), and on Kotlin/Native and
Android the formats ran on the priority-0 scalar floor. New `skainet_q5_0_matmul` /
`skainet_q5_1_matmul` C kernels (plain NEON, no dotprod/i8mm requirement — runs on every
AArch64 core) expand the `qh` high-bit plane with a per-lane `vtstq_u8` bit test and fold
the dequant algebraically (`d*(dot - 16*Σx)` for Q5_0, `d*dot + m*Σx` for Q5_1) so the
per-block input sum hoists out of the output-row loop. Wired into all three consumers:
the FFM `NativeKernelProvider` (JVM), the cinterop `NativeKnKernelProvider`
(Kotlin/Native), and the Android JNI bridge (`JniKernels.q50Matmul`/`q51Matmul` +
`JniKernelProvider`), each with parity tests against the scalar references. Unblocks the
packed Q5_1 path for `functiongemma-270m` "Q5_K_M" checkpoints (whose attention/FFN
weights are Q5_1) under `NATIVE_OPTIMIZED` — see SKaiNET-transformers#170. (#708)
### Performance

- **Primitive FP32 fast paths for the eager CPU ops** (`skainet-backend-cpu`,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@ set(SKAINET_KERNEL_SOURCES
${SKAINET_KERNELS_ROOT}/src/bf16_matmul.c
${SKAINET_KERNELS_ROOT}/src/fp16_matmul.c
${SKAINET_KERNELS_ROOT}/src/q4_0_matmul.c
${SKAINET_KERNELS_ROOT}/src/q5_0_matmul.c
${SKAINET_KERNELS_ROOT}/src/q5_1_matmul.c
${SKAINET_KERNELS_ROOT}/src/q8_0_matmul.c
${SKAINET_KERNELS_ROOT}/src/q4k_matmul.c
${SKAINET_KERNELS_ROOT}/src/q5k_matmul.c
Expand Down
28 changes: 28 additions & 0 deletions skainet-backends/skainet-backend-jni-cpu/native/skainet_jni.c
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,34 @@ Java_sk_ainet_exec_kernel_jni_JniKernels_q40Matmul(
inputDim, outputDim, out, outputOffset))
}

JNIEXPORT void JNICALL
Java_sk_ainet_exec_kernel_jni_JniKernels_q50Matmul(
JNIEnv* env, jobject thiz,
jfloatArray input, jint inputOffset,
jbyteArray weight, jint weightByteOffset,
jint inputDim, jint outputDim,
jfloatArray output, jint outputOffset
) {
(void) thiz;
SKAINET_JNI_MATMUL_BODY(
skainet_q5_0_matmul(in, inputOffset, (const uint8_t*) w, weightByteOffset,
inputDim, outputDim, out, outputOffset))
}

JNIEXPORT void JNICALL
Java_sk_ainet_exec_kernel_jni_JniKernels_q51Matmul(
JNIEnv* env, jobject thiz,
jfloatArray input, jint inputOffset,
jbyteArray weight, jint weightByteOffset,
jint inputDim, jint outputDim,
jfloatArray output, jint outputOffset
) {
(void) thiz;
SKAINET_JNI_MATMUL_BODY(
skainet_q5_1_matmul(in, inputOffset, (const uint8_t*) w, weightByteOffset,
inputDim, outputDim, out, outputOffset))
}

JNIEXPORT void JNICALL
Java_sk_ainet_exec_kernel_jni_JniKernels_q4kMatmul(
JNIEnv* env, jobject thiz,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@ import org.junit.Test
import org.junit.runner.RunWith
import sk.ainet.exec.kernel.ScalarQ4_0MatmulKernel
import sk.ainet.exec.kernel.ScalarQ4_KMatmulKernel
import sk.ainet.exec.kernel.ScalarQ5_0MatmulKernel
import sk.ainet.exec.kernel.ScalarQ5_1MatmulKernel
import sk.ainet.exec.kernel.ScalarQ6_KMatmulKernel
import sk.ainet.exec.kernel.ScalarQ8_0MatmulKernel
import kotlin.math.abs
Expand Down Expand Up @@ -100,6 +102,8 @@ class JniKernelParityTest {

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 conditionQ5_0: (ByteArray, Int) -> Unit = { b, base -> fp16One(b, base) }
private val conditionQ5_1: (ByteArray, Int) -> Unit = { b, base -> fp16One(b, base); fp16One(b, base + 2) }
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) }

Expand Down Expand Up @@ -181,6 +185,18 @@ class JniKernelParityTest {
reference = ScalarQ4_0MatmulKernel::matmul, jni = JniKernels::q40Matmul,
)

@Test
fun q50_parity() = assertExactParity(
1024, 64, 11, 2e-1f, 32, 22, conditionQ5_0,
reference = ScalarQ5_0MatmulKernel::matmul, jni = JniKernels::q50Matmul,
)

@Test
fun q51_parity() = assertExactParity(
1024, 64, 13, 2e-1f, 32, 24, conditionQ5_1,
reference = ScalarQ5_1MatmulKernel::matmul, jni = JniKernels::q51Matmul,
)

@Test
fun q4k_parity_single_block() = assertRmsParity(
256, 16, 42, 1e-2f, 256, 144, conditionQ4K,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@ import sk.ainet.backend.api.kernel.KernelProvider
import sk.ainet.backend.api.kernel.Q4KMatmulKernel
import sk.ainet.backend.api.kernel.Q4_0MatmulKernel
import sk.ainet.backend.api.kernel.Q5KMatmulKernel
import sk.ainet.backend.api.kernel.Q5_0MatmulKernel
import sk.ainet.backend.api.kernel.Q5_1MatmulKernel
import sk.ainet.backend.api.kernel.Q6KMatmulKernel
import sk.ainet.backend.api.kernel.Q8_0MatmulKernel
import sk.ainet.backend.api.kernel.Fp32MatmulKernel
Expand Down Expand Up @@ -55,6 +57,10 @@ public object JniKernelProvider : KernelProvider {

override fun matmulQ6K(): Q6KMatmulKernel? = if (available) JniQ6KMatmul else null

override fun matmulQ5_0(): Q5_0MatmulKernel? = if (available) JniQ5_0Matmul else null

override fun matmulQ5_1(): Q5_1MatmulKernel? = if (available) JniQ5_1Matmul else null

private object JniQ8_0Matmul : Q8_0MatmulKernel {
override fun matmul(
input: FloatArray, inputOffset: Int,
Expand All @@ -77,6 +83,28 @@ public object JniKernelProvider : KernelProvider {
)
}

private object JniQ5_0Matmul : Q5_0MatmulKernel {
override fun matmul(
input: FloatArray, inputOffset: Int,
weight: ByteArray, weightByteOffset: Int,
inputDim: Int, outputDim: Int,
output: FloatArray, outputOffset: Int,
): Unit = JniKernels.q50Matmul(
input, inputOffset, weight, weightByteOffset, inputDim, outputDim, output, outputOffset
)
}

private object JniQ5_1Matmul : Q5_1MatmulKernel {
override fun matmul(
input: FloatArray, inputOffset: Int,
weight: ByteArray, weightByteOffset: Int,
inputDim: Int, outputDim: Int,
output: FloatArray, outputOffset: Int,
): Unit = JniKernels.q51Matmul(
input, inputOffset, weight, weightByteOffset, inputDim, outputDim, output, outputOffset
)
}

private object JniQ4KMatmul : Q4KMatmulKernel {
override fun matmul(
input: FloatArray, inputOffset: Int,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,20 @@ public object JniKernels {
output: FloatArray, outputOffset: Int,
)

public external fun q50Matmul(
input: FloatArray, inputOffset: Int,
weight: ByteArray, weightByteOffset: Int,
inputDim: Int, outputDim: Int,
output: FloatArray, outputOffset: Int,
)

public external fun q51Matmul(
input: FloatArray, inputOffset: Int,
weight: ByteArray, weightByteOffset: Int,
inputDim: Int, outputDim: Int,
output: FloatArray, outputOffset: Int,
)

public external fun q4kMatmul(
input: FloatArray, inputOffset: Int,
weight: ByteArray, weightByteOffset: Int,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@ set(SKAINET_KERNEL_SOURCES
src/fp16_matmul.c
src/q8_0_matmul.c
src/q4_0_matmul.c
src/q5_0_matmul.c
src/q5_1_matmul.c
)

# SHARED: consumed by the JVM via java.lang.foreign (FFM), bundled as a JAR
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -214,6 +214,62 @@ SKAINET_API void skainet_q4_0_matmul(
int32_t output_offset
);

/*
* Q5_0 matrix-vector multiply.
*
* output[output_offset + o] = sum_j input[input_offset + j] *
* dequant(weight[block, o, j])
*
* Block layout: canonical ggml Q5_0, 32 elements per block, 22 bytes
* per block (2 B FP16 scale `d` + 4 B `qh` high-bit plane + 16 B packed
* 4-bit codes in split layout — low nibbles → elements 0..15, high
* nibbles → 16..31; bit j of the little-endian 32-bit `qh` is the fifth
* bit of element j), with packed weights laid out as
* weight + weight_byte_offset + (block_idx * output_dim + o) * 22
*
* Dequant per element: `(code - 16) * d` with
* `code = nibble | (fifth_bit << 4)`. input_dim must be a multiple
* of 32.
*/
SKAINET_API void skainet_q5_0_matmul(
const float* input,
int32_t input_offset,
const uint8_t* weight,
int32_t weight_byte_offset,
int32_t input_dim,
int32_t output_dim,
float* output,
int32_t output_offset
);

/*
* Q5_1 matrix-vector multiply.
*
* output[output_offset + o] = sum_j input[input_offset + j] *
* dequant(weight[block, o, j])
*
* Block layout: canonical ggml Q5_1, 32 elements per block, 24 bytes
* per block (2 B FP16 scale `d` + 2 B FP16 min `m` + 4 B `qh` high-bit
* plane + 16 B packed 4-bit codes in split layout; bit j of the
* little-endian 32-bit `qh` is the fifth bit of element j), with packed
* weights laid out as
* weight + weight_byte_offset + (block_idx * output_dim + o) * 24
*
* Dequant per element: `d * code + m` (affine, no re-centring) with
* `code = nibble | (fifth_bit << 4)`. input_dim must be a multiple
* of 32.
*/
SKAINET_API void skainet_q5_1_matmul(
const float* input,
int32_t input_offset,
const uint8_t* weight,
int32_t weight_byte_offset,
int32_t input_dim,
int32_t output_dim,
float* output,
int32_t output_offset
);

#ifdef __cplusplus
}
#endif
Expand Down
Loading
Loading