Skip to content

[bitnet] T2: lm_head via BITNET_PLANES weight format (WeightForm) + two-stage top-k decode #337

Description

@michalharakal

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

  • Weight form wiring (in the T1 BitNet loader): request EncodingRequest.RequantizeTo(BITNET_PLANES) for output.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.
  • Two-stage decode (NeoGPU Stage 2): take top-200 candidates from the Stage-1 logits (planes 0–3), refine only those rows with planes 4–7 (weights 1/3⁴..1/3⁷) via the codec's row accessors. Needs a candidate-rescoring API next to sampleFromLogits in SamplingUtils / GenerateExtensions — the current contract is a full-vocab logits array. Keep the single-stage path as default until this lands.

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

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions