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
28 changes: 28 additions & 0 deletions build-logic/convention/src/main/kotlin/GenerateKernelMatrixTask.kt
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,34 @@ abstract class GenerateKernelMatrixTask : DefaultTask() {
}
appendLine("|===")
appendLine("")
if (module.mapped.isNotEmpty()) {
appendLine("== Mapped serving (row-major, off-heap)")
appendLine("")
appendLine(
"Kernels that read the weight in canonical row-major GGUF file order straight " +
"from off-heap bytes (mmap'd pages or a direct buffer) — no heap staging, no " +
"prepack copy (#1189). This is what lets a model larger than the managed-heap " +
"cap decode on Android. Dense `Float32` tensors are also served from the " +
"mapping (as element views, not a matmul kernel) — see " +
"`StorageCapabilities.mappedServableEncodings` for the authoritative set the " +
"memory plan budgets against the page cache. An empty cell means the format " +
"heap-stages under `WeightResidency.MAPPED` on that platform.",
)
appendLine("")
appendLine("[cols=\"$colSpec\", options=\"header\"]")
appendLine("|===")
append("| Weight format ")
module.platforms.forEach { append("| $it ") }
appendLine("")
appendLine("")
module.mapped.forEach { fmt ->
append("| `${fmt.name}` ")
module.platforms.forEach { p -> append("| ${fmt.byPlatform[p] ?: "—"} ") }
appendLine("")
}
appendLine("|===")
appendLine("")
}
appendLine(
"See also the eager backends & kernels mindmap " +
"(xref:explanation/eager-execution.adoc[]) for the narrative overview and gaps.",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,12 @@ data class KernelSupportModule(
val inputDtype: String = "Float32",
val platforms: List<String> = emptyList(),
val formats: List<KernelFormatSupport> = emptyList(),
/**
* Mapped serving (#1189): formats whose weight a kernel reads in canonical row-major GGUF
* file order straight from off-heap (mmap'd/direct-buffer) bytes — no heap staging, no
* prepack. Absent/empty in pre-#1189 JSON, so old files stay decodable.
*/
val mapped: List<KernelFormatSupport> = emptyList(),
)

@Serializable
Expand Down
12 changes: 12 additions & 0 deletions docs/modules/ROOT/pages/reference/kernel-support-matrix.adoc
Original file line number Diff line number Diff line change
Expand Up @@ -20,4 +20,16 @@ Each cell is the best (highest-priority) provider that serves `Float32 × format
| `Q5_0` | native-ffm | native-jni | native-cinterop | native-cinterop | scalar
|===

== Mapped serving (row-major, off-heap)

Kernels that read the weight in canonical row-major GGUF file order straight from off-heap bytes (mmap'd pages or a direct buffer) — no heap staging, no prepack copy (#1189). This is what lets a model larger than the managed-heap cap decode on Android. Dense `Float32` tensors are also served from the mapping (as element views, not a matmul kernel) — see `StorageCapabilities.mappedServableEncodings` for the authoritative set the memory plan budgets against the page cache. An empty cell means the format heap-stages under `WeightResidency.MAPPED` on that platform.

[cols="1,1,1,1,1,1", options="header"]
|===
| Weight format | JVM | Android | Native·linux | Native·apple | JS/WASM

| `Q4_K` | — | native-jni-direct | — | — | —
| `Q6_K` | — | native-jni-direct | — | — | —
|===

See also the eager backends & kernels mindmap (xref:explanation/eager-execution.adoc[]) for the narrative overview and gaps.
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ set(SKAINET_KERNEL_SOURCES
${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/skainet_row_threads.c
${SKAINET_KERNELS_ROOT}/src/q4k_matmul.c
${SKAINET_KERNELS_ROOT}/src/q5k_matmul.c
${SKAINET_KERNELS_ROOT}/src/q6k_matmul.c
Expand Down
51 changes: 51 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 @@ -143,6 +143,57 @@ Java_sk_ainet_exec_kernel_jni_JniKernels_q6kMatmul(
inputDim, outputDim, out, outputOffset))
}

/*
* Direct-buffer row-major matmuls (#1189): the weight arrives as a direct
* ByteBuffer over mmap'd (or direct-allocated) bytes instead of a heap
* ByteArray, and stays in canonical GGUF row-major block order — which is
* what lets a model's quantized payloads never touch the managed heap.
*
* GetDirectBufferAddress is a JNI call, so it MUST run before the critical
* pins (no JNI calls are allowed between Get..Critical and Release..Critical).
* A NULL address (non-direct buffer) leaves the output untouched; the Kotlin
* caller guarantees directness by construction (DirectBufferStorage /
* MappedBufferStorage hand out direct buffers only).
*/
#define SKAINET_JNI_MATMUL_RM_DIRECT_BODY(CALL) \
const uint8_t* w = \
(const uint8_t*) (*env)->GetDirectBufferAddress(env, weight); \
jfloat* in = w ? (*env)->GetPrimitiveArrayCritical(env, input, NULL) : NULL; \
jfloat* out = in ? (*env)->GetPrimitiveArrayCritical(env, output, NULL) : NULL; \
if (out) { \
CALL; \
} \
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_q4kMatmulRmDirect(
JNIEnv* env, jobject thiz,
jfloatArray input, jint inputOffset,
jobject weight, jint weightByteOffset,
jint inputDim, jint outputDim,
jfloatArray output, jint outputOffset
) {
(void) thiz;
SKAINET_JNI_MATMUL_RM_DIRECT_BODY(
skainet_q4k_matmul_rm(in, inputOffset, w, weightByteOffset,
inputDim, outputDim, out, outputOffset))
}

JNIEXPORT void JNICALL
Java_sk_ainet_exec_kernel_jni_JniKernels_q6kMatmulRmDirect(
JNIEnv* env, jobject thiz,
jfloatArray input, jint inputOffset,
jobject weight, jint weightByteOffset,
jint inputDim, jint outputDim,
jfloatArray output, jint outputOffset
) {
(void) thiz;
SKAINET_JNI_MATMUL_RM_DIRECT_BODY(
skainet_q6k_matmul_rm(in, inputOffset, w, weightByteOffset,
inputDim, outputDim, out, outputOffset))
}

/*
* bitnet_gemv (SKEEP-003 §5.3, #1041): int8 activations against ternary TQ2_0
* weights. Its activation is a *byte* array, not floats, so it does not fit
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ import sk.ainet.backend.api.kernel.KernelDispatch
import sk.ainet.backend.api.kernel.KernelPacks
import sk.ainet.context.DirectCpuExecutionContext
import sk.ainet.exec.kernel.jni.JniKernelProvider
import sk.ainet.exec.kernel.jni.JniMappedKernelPack
import sk.ainet.io.MappedRandomAccessSource
import sk.ainet.io.gguf.StreamingGGUFReader
import sk.ainet.io.gguf.StreamingGgufParametersLoader
Expand Down Expand Up @@ -104,6 +105,10 @@ class M2A5DeviceMeasurement {
val prepack = args.getString("prepack")?.toBoolean() ?: false
val steps = arg("steps", 16)
val warmup = arg("warmup", 4)
// `residency=heap` (default mapped) restores heap staging — the #1193 A/B lever for
// models that fit the cap. The plan below prices the same form the load uses (#1190).
val residency = if (args.getString("residency") == "heap") WeightResidency.HEAP else WeightResidency.MAPPED
val loadForm = WeightForm(shape = WeightShapeOrientation.OUT_IN, residency = residency)

val report = StringBuilder()
fun line(s: String = "") { report.append(s).append('\n') }
Expand All @@ -115,13 +120,13 @@ class M2A5DeviceMeasurement {
line("- device: ${Build.MANUFACTURER} ${Build.MODEL}, Android ${Build.VERSION.RELEASE} (SDK ${Build.VERSION.SDK_INT}), ABI ${Build.SUPPORTED_ABIS.firstOrNull()}")
line("- ART heap cap (Runtime.maxMemory): ${mb(heapCap)}")
line("- model: ${modelFile.name}, ${mb(modelFile.length())} on disk")
line("- ctx=$ctxLen, decode steps=$steps (warm-up $warmup), prepack=$prepack")
line("- ctx=$ctxLen, decode steps=$steps (warm-up $warmup), prepack=$prepack, residency=${residency.name.lowercase()}")
line()

// ---- plan, from the header only --------------------------------------------------
val (plan, geometry) = MappedRandomAccessSource.open(modelPath).let { src ->
StreamingGGUFReader.open(src).use { reader ->
val input = reader.planInput(ctx = ctxLen)
val input = reader.planInput(ctx = ctxLen, formFor = { loadForm })
MemoryPlans.plan(input, Budget.of(heapCap)) to input.geometry
}
}
Expand All @@ -131,6 +136,10 @@ class M2A5DeviceMeasurement {
KernelDispatch.clearForTesting()
runCatching { KernelPacks.install(JniKernelProvider); KernelPacks.installPacked(JniKernelProvider) }
.onFailure { line("_note: JNI kernel pack unavailable (${it.message}); reference kernels serve — memory numbers stay valid, timings do not._") }
// #1189: row-major direct-buffer kernels — these serve the mapped packed weights below
// with zero copies and no prepack.
runCatching { JniMappedKernelPack.install() }
.onFailure { line("_note: mapped kernel pack unavailable (${it.message}); mapped packed weights fall back to the decoding reference._") }

val sink = RecordingTraceSink()
val ctx = DirectCpuExecutionContext()
Expand All @@ -140,28 +149,37 @@ class M2A5DeviceMeasurement {
runBlocking {
StreamingGgufParametersLoader(
sourceProvider = { MappedRandomAccessSource.open(modelPath) },
weightForm = WeightForm(shape = WeightShapeOrientation.OUT_IN, residency = WeightResidency.MAPPED),
weightForm = loadForm,
traceSink = sink,
).load<FP32, Float>(ctx, FP32::class) { name, t -> tensors[name] = t }
}
val loadMs = (System.nanoTime() - tLoad0) / 1_000_000
val sLoaded = MemoryProbe.sample()

// Account the materialized weights as model-scope allocations so plan-vs-actual sees
// them; the loader's own storages are not all sink-wired.
// them; the loader's own storages are not all sink-wired. #1189 note: physicalBytes
// (not packedData.size) — a mapped packed tensor keeps its bytes off-heap and has no
// heap ByteArray at all; those bytes are counted separately as mapped.
var weightHeapBytes = 0L
var weightMappedBytes = 0L
var mappedPackedCount = 0
var syntheticStorageId = -1_000_000L
for ((name, t) in tensors) {
val bytes = (t.data as? PackedBlockStorage)?.packedData?.size?.toLong()
?: (t.shape.volume.toLong() * 4L)
weightHeapBytes += bytes
val d = t.data
val bytes = (d as? PackedBlockStorage)?.physicalBytes ?: (t.shape.volume.toLong() * 4L)
if (d is sk.ainet.lang.tensor.data.BufferPackedTensorData) {
weightMappedBytes += bytes; mappedPackedCount += 1
} else {
weightHeapBytes += bytes
}
syntheticStorageId -= 1
sink.emit(TraceEvent.Allocation(syntheticStorageId, ScopeKind.MODEL, bytes, null, "m2a5:$name"))
}
line("## Load")
line()
line("- wall time: ${loadMs} ms for ${tensors.size} tensors")
line("- materialized bytes (heap tensors + packed payloads): ${mb(weightHeapBytes)}")
line("- heap bytes (dense tensors + heap packed payloads): ${mb(weightHeapBytes)}")
line("- mapped packed payloads (#1189, off-heap): ${mb(weightMappedBytes)} in $mappedPackedCount tensors")
line("- RSS before → after load: ${mb(sBefore.rssBytes)} → ${mb(sLoaded.rssBytes)} (Δ ${mb((sLoaded.rssBytes ?: 0L) - (sBefore.rssBytes ?: 0L))})")
line("- major faults during load: ${sLoaded.majorFaultsSince(sBefore) ?: "—"}")
line()
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,121 @@
package sk.ainet.exec.kernel.jni

import java.nio.ByteBuffer
import sk.ainet.backend.api.kernel.KernelDispatch
import sk.ainet.backend.api.kernel.KernelKey
import sk.ainet.backend.api.kernel.LayoutClass
import sk.ainet.backend.api.kernel.OperandKey
import sk.ainet.backend.api.kernel.ReferenceMatmulKernel
import sk.ainet.backend.api.kernel.ViewKernel
import sk.ainet.lang.memory.BlockOrder
import sk.ainet.lang.memory.DirectBufferStorage
import sk.ainet.lang.memory.ExperimentalMemoryApi
import sk.ainet.lang.memory.Format
import sk.ainet.lang.memory.MappedBufferStorage
import sk.ainet.lang.memory.Storage
import sk.ainet.lang.memory.TensorView
import sk.ainet.lang.tensor.storage.TensorEncoding
import sk.ainet.lang.types.FP32

/**
* A [ViewKernel] over the JNI **direct-buffer row-major** matmuls (#1189): the weight is a
* `BLOCKED_ROW_MAJOR` packed view whose storage is off-heap ([MappedBufferStorage] — mmap'd GGUF
* pages — or [DirectBufferStorage]), read by the native kernel straight at its address. No heap
* `ByteArray`, no relayout: this is what lets a model bigger than the ART cap decode on Android.
*
* Contrast with [sk.ainet.backend.api.kernel.PackedViewMatmulKernel], which serves the same
* encodings from heap bytes in `BLOCKED_INPUT_MAJOR` (prepacked feed) order. Registering both
* lets the dispatcher pick by what the weight actually is — canonical mapped weights match this
* key exactly and are served with zero copies.
*/
@ExperimentalMemoryApi
public class JniBufferPackedMatmulKernel(
encodingName: String,
override val key: KernelKey,
private val matmul: (
input: FloatArray, inputOffset: Int,
weight: ByteBuffer, weightByteOffset: Int,
inputDim: Int, outputDim: Int,
output: FloatArray, outputOffset: Int,
) -> Unit,
) : ViewKernel {

override val name: String = "jni-buffer-$encodingName"

override fun run(inputs: List<TensorView>, out: TensorView) {
require(inputs.size == 2) { "matmul takes two operands" }
val a = inputs[0]
val w = inputs[1]
val rows = a.shape[0]
val k = a.shape[1]
val n = w.shape[0]
require(w.shape[1] == k) { "inner dimensions disagree: [$rows, $k] × [$n, ${w.shape[1]}]" }
check(w.layout.blockOrder == BlockOrder.ROW_MAJOR) {
"$name reads canonical row-major weights; this one is ${w.layout.blockOrder}"
}

val buffer = when (val s = w.storage) {
is MappedBufferStorage -> s.buffer()
is DirectBufferStorage -> s.buffer()
else -> return fallback(inputs, out)
}
val aHeap = a.storage as? Storage.Heap ?: return fallback(inputs, out)
val oHeap = out.storage as? Storage.Heap ?: return fallback(inputs, out)
val activation = aHeap.floats ?: return fallback(inputs, out)
val output = oHeap.floats ?: return fallback(inputs, out)
// The JNI kernel takes a contiguous activation row; a strided one would be mis-indexed.
if (!a.isContiguous) return fallback(inputs, out)

val weightOffset = (w.layout.offsetElements * w.layout.elementBytes).toInt()
for (r in 0 until rows) {
matmul(
activation, aHeap.arrayOffset + (a.layout.offsetElements + r.toLong() * k).toInt(),
buffer, weightOffset,
k, n,
output, oHeap.arrayOffset + (out.layout.offsetElements + r.toLong() * n).toInt(),
)
}
}

private fun fallback(inputs: List<TensorView>, out: TensorView) {
ReferenceMatmulKernel(key).run(inputs, out)
}

public companion object {
/** The key this kernel serves: dense contiguous FP32 activation × row-major packed weight. */
public fun keyFor(encoding: TensorEncoding): KernelKey = KernelKey(
op = "matmul",
operands = listOf(
OperandKey.contiguous(Format.dense(FP32)),
OperandKey(Format(FP32, encoding), LayoutClass.BLOCKED_ROW_MAJOR),
),
)
}
}

/**
* Registers the direct-buffer row-major kernels (#1189) into [KernelDispatch]. Call next to
* `KernelPacks.install(JniKernelProvider)` on Android when weights are loaded with
* `WeightResidency.MAPPED` — without it a mapped packed weight falls back to the decoding
* reference kernel (correct, hours-scale slow).
*/
@ExperimentalMemoryApi
public object JniMappedKernelPack {
public fun install() {
if (!JniKernelProvider.isAvailable()) return
KernelDispatch.register(
JniBufferPackedMatmulKernel(
TensorEncoding.Q4_K.name,
JniBufferPackedMatmulKernel.keyFor(TensorEncoding.Q4_K),
JniKernels::q4kMatmulRmDirect,
),
)
KernelDispatch.register(
JniBufferPackedMatmulKernel(
TensorEncoding.Q6_K.name,
JniBufferPackedMatmulKernel.keyFor(TensorEncoding.Q6_K),
JniKernels::q6kMatmulRmDirect,
),
)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -171,4 +171,25 @@ public object JniKernels {
inputDim: Int, outputDim: Int,
output: FloatArray, outputOffset: Int,
)

/**
* Row-major Q4_K matmul over a **direct** [java.nio.ByteBuffer] weight (#1189): the blocks
* stay in canonical GGUF file order — `(o * blocksPerRow + b) * 144` from [weightByteOffset]
* — and are read straight out of the buffer's off-heap (typically mmap'd) memory. The buffer
* MUST be direct; a heap buffer has no native address and the call writes nothing.
*/
public external fun q4kMatmulRmDirect(
input: FloatArray, inputOffset: Int,
weight: java.nio.ByteBuffer, weightByteOffset: Int,
inputDim: Int, outputDim: Int,
output: FloatArray, outputOffset: Int,
)

/** Row-major Q6_K matmul over a direct ByteBuffer weight (#1189); see [q4kMatmulRmDirect]. */
public external fun q6kMatmulRmDirect(
input: FloatArray, inputOffset: Int,
weight: java.nio.ByteBuffer, weightByteOffset: Int,
inputDim: Int, outputDim: Int,
output: FloatArray, outputOffset: Int,
)
}
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ endif()
set(SKAINET_KERNEL_SOURCES
src/skainet_smoke.c
src/skainet_cpu_features.c
src/skainet_row_threads.c
src/q4k_matmul.c
src/q5k_matmul.c
src/q6k_matmul.c
Expand Down
Loading
Loading