Part of the BitNet tracking issue (T2, depends on T1 and on SKaiNET's BITNET_PLANES encoding — SKaiNET-developers/SKaiNET#1136 Phase 6).
Approach — no new op, layer, or fusion pass
This deliberately stays inside standard SKaiNET physiology. The network def keeps its stock VoidDense("output", vocabSize, dim); the DAG keeps its ordinary matmul(x, transpose(output.weight)) node. The fused kernel is reached purely by weight format: the lm_head weight is loaded as TensorEncoding.BITNET_PLANES, and KernelDispatch's exact-key match routes the matmul to the skainet_ternary_lmhead_stage1-backed kernel. DIRECT mode and the compiled/optimized path both funnel through ops.matmul, so both get the fast path with zero per-mode hooks. (Earlier draft's LLMFusionPass rule / fused_lm_head GenericOperation / custom TernaryLmHead module: dropped.)
What remains here is genuinely application-level: choosing the format for the output weight, and the two-stage sampling algorithm.
Tasks
Tests
- Logits parity on a real BitNet GGUF:
BITNET_PLANES lm_head vs exact BITNET_B1_58 lm_head — Stage-1-only is an approximation (plane truncation ≈ 1.85% of row scale); assert top-k stability rather than bitwise equality, and that two-stage rescoring restores the exact top-1
- Decode-mode coverage: DIRECT and OPTIMIZED produce identical logits (both dispatch by format, no mode-specific code)
- Prefill unaffected (
forwardBatched's full-vocab-then-narrow waste is a separate follow-up issue)
Notes
- Upstream driver reference: NeoGPU
src/hs_ml_infer.c:1063–1175 (lmh_plane_w[8], LMH_CANDIDATES = 200)
hs_ml_lmhead_stage1 always spawns 4 pthreads (no threshold) — appropriate at vocab scale
Part of the BitNet tracking issue (T2, depends on T1 and on SKaiNET's
BITNET_PLANESencoding — SKaiNET-developers/SKaiNET#1136 Phase 6).Approach — no new op, layer, or fusion pass
This deliberately stays inside standard SKaiNET physiology. The network def keeps its stock
VoidDense("output", vocabSize, dim); the DAG keeps its ordinarymatmul(x, transpose(output.weight))node. The fused kernel is reached purely by weight format: the lm_head weight is loaded asTensorEncoding.BITNET_PLANES, andKernelDispatch's exact-key match routes the matmul to theskainet_ternary_lmhead_stage1-backed kernel. DIRECT mode and the compiled/optimized path both funnel throughops.matmul, so both get the fast path with zero per-mode hooks. (Earlier draft'sLLMFusionPassrule /fused_lm_headGenericOperation/ customTernaryLmHeadmodule: dropped.)What remains here is genuinely application-level: choosing the format for the output weight, and the two-stage sampling algorithm.
Tasks
EncodingRequest.RequantizeTo(BITNET_PLANES)foroutput.weight— the plane encoding runs once at load through SKaiNET's WeightForm machinery, traced like any other conversion. Make it a loader option (default on for BitNet), so the exact single-plane path remains selectable.sampleFromLogitsinSamplingUtils/GenerateExtensions— the current contract is a full-vocab logits array. Keep the single-stage path as default until this lands.Tests
BITNET_PLANESlm_head vs exactBITNET_B1_58lm_head — Stage-1-only is an approximation (plane truncation ≈ 1.85% of row scale); assert top-k stability rather than bitwise equality, and that two-stage rescoring restores the exact top-1forwardBatched's full-vocab-then-narrow waste is a separate follow-up issue)Notes
src/hs_ml_infer.c:1063–1175(lmh_plane_w[8],LMH_CANDIDATES = 200)hs_ml_lmhead_stage1always spawns 4 pthreads (no threshold) — appropriate at vocab scale