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
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -50,3 +50,4 @@ out/

### Build artifacts ###
com/
.cxx/
16 changes: 16 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,22 @@

## [Unreleased]

### Added

- **`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
sources via NDK/CMake with thin JNI shims (`GetPrimitiveArrayCritical`, zero-copy pins) and
a priority-100 `JniKernelProvider` discovered via `ServiceLoader` (which ART supports; the
Android ops factory now installs discovered providers exactly like the JVM does). Two `.so`
tiers are built from the same sources and selected at load time from `/proc/cpuinfo`:
baseline `armv8-a` (NEON, runs on every arm64 core — 0 dot-product instructions, verified by
disassembly) and `armv8.2-a+fp16+dotprod` (enables the `vdotq_s32` q4k/q6k paths — would
SIGILL on Cortex-A53-class cores, hence the gate). Q8_0/Q4_0/Q4_K/Q5_K/Q6_K bridged;
16 KB-page-aligned `.so`s (Android 15+); consumer R8 rules keep the ServiceLoader entry in
release builds. On-device parity tests included (`androidTest`, vs the scalar references).
Part of the mobile-kernels effort (#920).

### Fixed

- **`TensorData.copyToFloatArray()` default implementation works for rank >= 2.** It used to
Expand Down
2 changes: 2 additions & 0 deletions gradle/libs.versions.toml
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,8 @@ npm-brace-expansion = "2.1.4" # GHSA-rgw5-rvv9-x895, GHSA-mh99-v99m-4gvg
jackson-databind = { module = "com.fasterxml.jackson.core:jackson-databind", version.ref = "jacksonDatabind" }
json-schema-validator = { module = "com.networknt:json-schema-validator", version.ref = "jsonSchemaValidator" }
junit = { module = "junit:junit", version.ref = "junit" }
androidx-test-junit = { module = "androidx.test.ext:junit", version = "1.2.1" }
androidx-test-runner = { module = "androidx.test:runner", version = "1.6.2" }
junit-jupiter = { module = "org.junit.jupiter:junit-jupiter", version.ref = "junitJupiter" }
junit-platform-launcher = { module = "org.junit.platform:junit-platform-launcher" }
kotlinx-coroutines = { group = "org.jetbrains.kotlinx", name = "kotlinx-coroutines-core", version.ref = "kotlinxCoroutines" }
Expand Down
1 change: 1 addition & 0 deletions settings.gradle.kts
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ include("skainet-compile:skainet-compile-minerva")
include("skainet-backends:skainet-backend-api")
include("skainet-backends:skainet-backend-cpu")
include("skainet-backends:skainet-backend-native-cpu")
include("skainet-backends:skainet-backend-jni-cpu")

// ====== BENCHMARKS
include("skainet-backends:benchmarks:jvm-cpu-jmh")
Expand Down
Original file line number Diff line number Diff line change
@@ -1,13 +1,25 @@
package sk.ainet.exec.tensor.ops

import java.util.ServiceLoader
import sk.ainet.backend.api.kernel.KernelProvider
import sk.ainet.backend.api.kernel.KernelRegistry
import sk.ainet.exec.kernel.ScalarKernelProvider
import sk.ainet.lang.tensor.data.TensorDataFactory
import sk.ainet.lang.tensor.ops.TensorOps

internal actual fun platformDefaultCpuOpsFactory(): (TensorDataFactory) -> TensorOps {
// Non-JVM has no ServiceLoader; register the scalar packed-quant kernels
// (Q4_K/Q6_K/Q5_1/Q5_0/Q8_0/Q4_0) so DefaultCpuOpsBase can dispatch them.
// ART supports java.util.ServiceLoader, so Android discovers kernel
// providers the same way the JVM does (#920): modules like
// skainet-backend-jni-cpu ship a META-INF/services entry and are
// picked up here at priority 100. Discovery failures are non-fatal —
// a provider that can't load just doesn't register.
runCatching {
ServiceLoader.load(KernelProvider::class.java).forEach { provider ->
runCatching { KernelRegistry.register(provider) }
}
}
// Scalar reference last: priority 0, always available — the floor the
// registry cascades to when no accelerated provider carries a kernel.
KernelRegistry.register(ScalarKernelProvider)
return { factory -> DefaultCpuOps(factory) }
}
75 changes: 75 additions & 0 deletions skainet-backends/skainet-backend-jni-cpu/build.gradle.kts
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
plugins {
// AGP 9 ships built-in Kotlin support — applying kotlin("android") is an
// error since 9.0, so this module uses the android plugin alone.
alias(libs.plugins.androidLibrary)
alias(libs.plugins.vanniktech.mavenPublish)
}

/*
* Android JNI bridge for the hand-written C/NEON matmul kernels
* (skainet-backend-native-cpu/native). ART has no java.lang.foreign, so the
* FFM provider can never run on Android — this module ships the same C
* sources as an AAR (.so via externalNativeBuild/NDK) with thin JNI shims
* and a priority-100 KernelProvider discovered via ServiceLoader (#920).
*
* Two shared libs are built from the same sources (see native/CMakeLists.txt):
* libskainet_jni.so — baseline armv8-a (NEON always present on arm64;
* runs on every device incl. Cortex-A53)
* libskainet_jni_v82.so — -march=armv8.2-a+fp16+dotprod (vdotq_s32 paths
* in q4k/q6k; would SIGILL on armv8.0 cores)
* JniKernels picks ONE at load time from /proc/cpuinfo features — runtime
* dispatch without symbol renaming or ifunc.
*/
android {
namespace = "sk.ainet.exec.kernel.jni"
compileSdk = libs.versions.android.compileSdk.get().toInt()

defaultConfig {
minSdk = libs.versions.android.minSdk.get().toInt()
consumerProguardFiles("consumer-rules.pro")
ndk {
// arm64 is the target that matters; x86_64 keeps emulator CI and
// desktop AVDs working (scalar C paths). 32-bit ARM is out of
// scope — different intrinsics story, shrinking device share.
abiFilters += listOf("arm64-v8a", "x86_64")
}
externalNativeBuild {
cmake {
arguments += "-DANDROID_STL=none"
}
}
testInstrumentationRunner = "androidx.test.runner.AndroidJUnitRunner"
}

externalNativeBuild {
cmake {
path = file("native/CMakeLists.txt")
}
}

compileOptions {
sourceCompatibility = JavaVersion.VERSION_11
targetCompatibility = JavaVersion.VERSION_11
}

kotlin {
compilerOptions {
jvmTarget.set(org.jetbrains.kotlin.gradle.dsl.JvmTarget.JVM_11)
}
}

// Publishing variant selection is configured by the vanniktech
// maven-publish plugin (AndroidSingleVariantLibrary "release").
}

dependencies {
api(project(":skainet-backends:skainet-backend-api"))

testImplementation(libs.junit)
testImplementation(libs.kotlin.test)

androidTestImplementation(libs.androidx.test.junit)
androidTestImplementation(libs.androidx.test.runner)
// Scalar reference kernels for on-device parity checks.
androidTestImplementation(project(":skainet-backends:skainet-backend-cpu"))
}
15 changes: 15 additions & 0 deletions skainet-backends/skainet-backend-jni-cpu/consumer-rules.pro
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
# ServiceLoader discovery of the JNI kernel provider (#920).
#
# R8 full mode strips classes that are only referenced from
# META-INF/services files. Without these rules a RELEASE build silently
# loses the provider and inference falls back to the 100x-slower scalar
# path with no error — the exact failure mode these rules exist to prevent.
-keep class sk.ainet.exec.kernel.jni.JniKernelProviderFactory { *; }
-keep class sk.ainet.exec.kernel.jni.JniKernels { *; }

# Native method registration: JNI resolves Java_sk_ainet_exec_kernel_jni_*
# symbols against these exact names — neither class nor method names may be
# renamed or removed.
-keepclasseswithmembers class sk.ainet.exec.kernel.jni.JniKernels {
native <methods>;
}
2 changes: 2 additions & 0 deletions skainet-backends/skainet-backend-jni-cpu/gradle.properties
Original file line number Diff line number Diff line change
@@ -0,0 +1,2 @@
POM_ARTIFACT_ID=skainet-backend-jni-cpu
POM_NAME=skainet Android (JNI) CPU kernel provider
51 changes: 51 additions & 0 deletions skainet-backends/skainet-backend-jni-cpu/native/CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
cmake_minimum_required(VERSION 3.22)
project(skainet_jni C)

# The C kernel sources are shared with skainet-backend-native-cpu — this
# module adds only the JNI shim. Single source of truth for the math.
set(SKAINET_KERNELS_ROOT ${CMAKE_CURRENT_SOURCE_DIR}/../../skainet-backend-native-cpu/native)

set(SKAINET_KERNEL_SOURCES
${SKAINET_KERNELS_ROOT}/src/skainet_smoke.c
${SKAINET_KERNELS_ROOT}/src/fp32_matmul.c
${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/q8_0_matmul.c
${SKAINET_KERNELS_ROOT}/src/q4k_matmul.c
${SKAINET_KERNELS_ROOT}/src/q5k_matmul.c
${SKAINET_KERNELS_ROOT}/src/q6k_matmul.c
)

set(SKAINET_JNI_SHIM ${CMAKE_CURRENT_SOURCE_DIR}/skainet_jni.c)

# Optimization flags match the sibling CMakeLists (-O3 -ffast-math) so the
# JNI path performs identically to the FFM/cinterop consumers of the same C.
set(SKAINET_C_FLAGS -O3 -ffast-math -funroll-loops)

# Android 15+ devices ship 16 KB pages; .so files must be max-page-size
# aligned. Recent NDKs default to this, the explicit flag pins it.
set(SKAINET_LINK_FLAGS -Wl,-z,max-page-size=16384)

function(skainet_add_jni_lib target)
add_library(${target} SHARED ${SKAINET_JNI_SHIM} ${SKAINET_KERNEL_SOURCES})
target_include_directories(${target} PRIVATE ${SKAINET_KERNELS_ROOT}/include)
target_compile_options(${target} PRIVATE ${SKAINET_C_FLAGS})
target_link_options(${target} PRIVATE ${SKAINET_LINK_FLAGS})
endfunction()

# Baseline: NDK default -march (armv8-a on arm64-v8a). NEON is architecturally
# guaranteed on AArch64, so fp32/q8_0/q4_0/q5k get their SIMD bodies, while the
# dotprod-guarded q4k/q6k paths fall back to scalar. Runs on EVERY arm64 core,
# including armv8.0 (Cortex-A53/A55-r0) — no SIGILL possible.
skainet_add_jni_lib(skainet_jni)

# v8.2 variant: enables the vdotq_s32 paths in q4k/q6k (and fp16 storage ops).
# Would SIGILL on armv8.0 cores — JniKernels loads it only after confirming
# `asimddp` (+ `asimdhp`/`fphp`) in /proc/cpuinfo. On non-ARM ABIs (x86_64
# emulator) the flags don't apply and the two libs are identical; both are
# still packaged so the loader's two-tier logic needs no per-ABI special case.
skainet_add_jni_lib(skainet_jni_v82)
if(ANDROID_ABI STREQUAL "arm64-v8a")
target_compile_options(skainet_jni_v82 PRIVATE -march=armv8.2-a+fp16+dotprod)
endif()
116 changes: 116 additions & 0 deletions skainet-backends/skainet-backend-jni-cpu/native/skainet_jni.c
Original file line number Diff line number Diff line change
@@ -0,0 +1,116 @@
#include <jni.h>
#include <stdint.h>

#include "skainet_kernels.h"

/*
* Thin JNI shims over the shared C matmul kernels (skainet_kernels.h).
* One function per kernel, no logic beyond array pinning + the call.
*
* Array strategy: GetPrimitiveArrayCritical for every array. On ART, heap
* primitive arrays are contiguous, so criticals are zero-copy pins. Rules
* honored here: no JNI calls between Get and Release, releases in reverse
* acquisition order, read-only arrays released with JNI_ABORT (no
* write-back), the output with 0 (write-back + unpin). Kernel calls are
* millisecond-scale — well within acceptable critical-section length.
*
* Method names deliberately contain no underscores (JNI mangles `_` to
* `_1`, which is easy to get wrong silently).
*/

/* Pin helper: acquire all three arrays, run `CALL`, release in reverse. */
#define SKAINET_JNI_MATMUL_BODY(CALL) \
jfloat* in = (*env)->GetPrimitiveArrayCritical(env, input, NULL); \
jbyte* w = in ? (*env)->GetPrimitiveArrayCritical(env, weight, NULL) : NULL; \
jfloat* out = w ? (*env)->GetPrimitiveArrayCritical(env, output, NULL) : NULL; \
if (out) { \
CALL; \
} \
if (out) (*env)->ReleasePrimitiveArrayCritical(env, output, out, 0); \
if (w) (*env)->ReleasePrimitiveArrayCritical(env, weight, w, JNI_ABORT); \
if (in) (*env)->ReleasePrimitiveArrayCritical(env, input, in, JNI_ABORT);

JNIEXPORT void JNICALL
Java_sk_ainet_exec_kernel_jni_JniKernels_smoke(
JNIEnv* env, jobject thiz,
jfloatArray input, jfloatArray output, jint length
) {
(void) thiz;
jfloat* in = (*env)->GetPrimitiveArrayCritical(env, input, NULL);
jfloat* out = in ? (*env)->GetPrimitiveArrayCritical(env, output, NULL) : NULL;
if (out) {
skainet_smoke_double(in, out, length);
}
if (out) (*env)->ReleasePrimitiveArrayCritical(env, output, out, 0);
if (in) (*env)->ReleasePrimitiveArrayCritical(env, input, in, JNI_ABORT);
}

JNIEXPORT void JNICALL
Java_sk_ainet_exec_kernel_jni_JniKernels_q80Matmul(
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_q8_0_matmul(in, inputOffset, (const uint8_t*) w, weightByteOffset,
inputDim, outputDim, out, outputOffset))
}

JNIEXPORT void JNICALL
Java_sk_ainet_exec_kernel_jni_JniKernels_q40Matmul(
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_q4_0_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,
jfloatArray input, jint inputOffset,
jbyteArray weight, jint weightByteOffset,
jint inputDim, jint outputDim,
jfloatArray output, jint outputOffset
) {
(void) thiz;
SKAINET_JNI_MATMUL_BODY(
skainet_q4k_matmul(in, inputOffset, (const uint8_t*) w, weightByteOffset,
inputDim, outputDim, out, outputOffset))
}

JNIEXPORT void JNICALL
Java_sk_ainet_exec_kernel_jni_JniKernels_q5kMatmul(
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_q5k_matmul(in, inputOffset, (const uint8_t*) w, weightByteOffset,
inputDim, outputDim, out, outputOffset))
}

JNIEXPORT void JNICALL
Java_sk_ainet_exec_kernel_jni_JniKernels_q6kMatmul(
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_q6k_matmul(in, inputOffset, (const uint8_t*) w, weightByteOffset,
inputDim, outputDim, out, outputOffset))
}
Loading
Loading