Skip to content

feat(backends): JNI + Kotlin/Native faces of the fused lm_head kernel (#1150) - #1172

Merged
michalharakal merged 2 commits into
developfrom
feature/1150-lmhead-jni-kn
Aug 26, 2026
Merged

michalharakal merged 2 commits into
developfrom
feature/1150-lmhead-jni-kn

Conversation

@michalharakal

Copy link
Copy Markdown
Contributor

Completes #1150 (follows #1171, which landed the encoding, codec, pack, FFM face, and loader requantizer).

What

  • JNI: ternaryLmheadStage1 shim (pins three arrays; derives the uint16_t* row-scale pointer from the pinned weight at rowScaleByteOffset), JniKernels external, JniTernaryLmhead : TernaryLmheadNative + install(). Baseline .so carries the full NEON path — no capability split.
  • Kotlin/Native: NativeKnTernaryLmhead over the auto-exposed cinterop symbol (.def unchanged); row-scale pointer = pinned weight addressOf(rowScaleByteOffset); inputDim == 0 zeroes the output without pinning empty arrays.
  • Parity tests on both bridges against a local reimplementation of the 4-plane fused contract over real codec-encoded BITNET_PLANES buffers — both plane windows (0–3 and 4–7, the two calls NativeTernaryPlanesViewKernel combines), BitNet hidden size (k=2560), and the kernel's internal 4-thread pool.

Verified locally

  • :skainet-backend-native-cpu:macosArm64Test — 3/3 green (real NEON through cinterop)
  • :skainet-backend-jni-cpu:assembleDebug + assembleDebugAndroidTest — both .so variants + test APK build with NDK
  • CI: linuxArm64Test -PcrossArm64=true qemu lane; JniKernelParityTest on an arm64 device

With this, every SKaiNET-side phase of #1136 is code-complete: #1137–#1141, #1150. Remaining: the #1166 compliance gate, then a release for transformers#335.

🤖 Generated with Claude Code

michalharakal and others added 2 commits August 26, 2026 15:41
JniTernaryLmhead and NativeKnTernaryLmhead complete the TernaryLmheadNative
seam over skainet_ternary_lmhead_stage1 — the fused 4-plane BITNET_PLANES
kernel now reaches all three bridges, closing #1150. The FP16 row scales
live inside the weight buffer, so both shims derive the uint16_t* from the
pinned weight at rowScaleByteOffset (2-byte aligned by the seam's
contract). No capability split, as with the gemv: baseline NEON only.

Parity on both bridges against a local reimplementation of the 4-plane
fused contract over real codec-encoded buffers — both plane windows (0-3
and 4-7, the two calls the view kernel combines), BitNet hidden size, and
the C kernel's internal 4-thread pool. Verified locally: macosArm64Test
green (real NEON through cinterop), JNI AAR + test APK build with NDK;
the qemu lane covers linuxArm64 in CI.

Closes #1150

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
…heck

The binary-compatibility validator's jvmApiCheck wants the .api dump to
name what #1171 made public — TernaryCodec's plane codec functions,
BitNetPlanesTensorData, and TensorEncoding.BITNET_PLANES. This also heals
the check on develop, which went red when #1171 merged without the dump.

Refs #1150

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
@michalharakal
michalharakal merged commit 2822617 into develop Aug 26, 2026
1 of 10 checks passed
@michalharakal
michalharakal deleted the feature/1150-lmhead-jni-kn branch August 26, 2026 15:47
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant