Skip to content

[bitnet] Wire BitNetTwoStageDecode into the generation loop (pre-lm_head hidden-state hook) #358

Description

@michalharakal

Problem

BitNetTwoStageDecode (stage-1 planes-0–3 scan + exact candidate rescoring with a provable top-k-equality bound) and the sampling API sampleFromCandidates/ScoredToken (llm-core/.../SamplingUtils.kt:154-190) shipped with PR #353 — but no generation loop calls them. OptimizedLLMRuntime exposes no pre-lm_head hidden state, and the CLI's BitNet branch decodes single-stage full-vocab. The two-stage path is reachable only from tests and library callers.

This is the residual #337 flagged as "a separate, small design": the hook was never added.

Proposed change

  • Runtime: expose the last-position hidden state before the lm_head projection (hook or return-shape addition in OptimizedLLMRuntime), so a family can substitute candidate scoring for the full-vocab matmul. Keep it format-driven — the hook is generic; only the BitNet path uses it today.
  • CLI: on the BitNet branch, use two-stage decode when the head weight is BITNET_PLANES (with [bitnet] Tied lm_head never engages the planes kernel on 2B4T — FP32-widened head, ~0.45 tok/s #357 landed this covers 2B4T's tied head). Single-stage remains the default everywhere else.

Verification

Refs: #335, #337 (residual), #357 (tied head), SKaiNET-developers/SKaiNET#1150.

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