Skip to content

feat(bitnet): wire the two-stage lm_head decode into generation (#358) - #362

Merged
michalharakal merged 1 commit into
developfrom
feat/bitnet-two-stage-decode
Aug 31, 2026
Merged

michalharakal merged 1 commit into
developfrom
feat/bitnet-two-stage-decode

Conversation

@michalharakal

Copy link
Copy Markdown
Contributor

Closes #358. Part of the BitNet tracking issue #335. Stacked on #361 (the tied planes head is what makes this engage on 2B4T) — merge #361 first; this PR then retargets to develop.

What

BitNetTwoStageDecode and sampleFromCandidates shipped with PR #353 but nothing called them. This adds the pre-lm_head hook #337 flagged as "a separate, small design" and drives NeoGPU's decode design end to end:

  • OptimizedLLMRuntime.forwardHidden(token) — runs the module-tree trunk (everything but the final lm_head projection) inside the same per-step forward scope; KV cache and position advance exactly as forward(). DIRECT/HYBRID only; OPTIMIZED compiles the full graph and has no partial execution to offer.
  • BitNetStage1Kernel (fun interface) + optional native stage-1 in BitNetTwoStageDecode — one fused lmhead_stage1 call scores all rows with planes 0–3; the exact planes matmul makes two calls to the same symbol, so the stage-1 scan halves the kernel work per step. The contract is restated locally rather than importing TernaryLmheadNative, because skainet-backend-api does not publish every native target this module builds for at 0.51 (androidNativeArm32 arrives with the 0.52 consumption).
  • generateTwoStage(...) — trunk-only forwards + bound-guaranteed exact top-200 rescoring + sampleFromCandidates. Prompt ingestion also runs forwardHidden: prefill logits were discarded anyway, so every prompt position saves its full-vocab projection outright.
  • CLI — enabled automatically on the BitNet path when the head loaded as BITNET_PLANES; prints "Two-stage lm_head decode: enabled" at load.

Measured (real microsoft/BitNet-b1.58-2B-4T I2_S, skainet-cli -s 24, macOS Apple Silicon, FFM kernels)

decode path tok/s
as-stored F16 tied head (pre-#357) 4.15
BITNET_PLANES head, exact matmul (#361) 4.54–4.72
two-stage (this PR) 6.20–6.29

Greedy decode stays coherent ("The capital of France is Paris. Paris is the capital of France. …") — and greedy is provably identical to the full matmul: topK's per-row bound rescores every row that could reach the top-k.

Tests

  • Greedy two-stage decode is token-for-token equal to the full-matmul loop on a tied synthetic GGUF — with the portable stage-1 and with the fused native stage-1.
  • Native stage-1 matches the Kotlin reference row-for-row.
  • jvm API dumps regenerated (llm-core's also picks up a stale generateUntilStop overload entry that predates this change; the legacy api/android file has no dump task in this configuration and is left as-is).

🤖 Generated with Claude Code

https://claude.ai/code/session_01YKeDSK4JF295y53Uvez954

BitNetTwoStageDecode and sampleFromCandidates shipped with the fold-in but
no generation loop called them. This adds the pre-lm_head hook #337 flagged
and drives NeoGPU's decode design end to end:

- OptimizedLLMRuntime.forwardHidden(token): runs the module-tree trunk —
  everything but the final lm_head projection — inside the same per-step
  forward scope; KV cache and position advance exactly as forward().
  DIRECT/HYBRID only (OPTIMIZED compiles the full graph).
- BitNetStage1Kernel (fun interface) + optional native stage-1 in
  BitNetTwoStageDecode: one fused lmhead_stage1 call scores all rows with
  planes 0-3 — the exact planes matmul makes two calls to the same symbol,
  so a stage-1 scan halves the kernel work per step. Restated locally
  instead of depending on backend-api, which does not publish every native
  target this module builds for at 0.51.
- generateTwoStage(prompt, steps, temperature, head, native, candidates):
  trunk-only forwards + bound-guaranteed top-200 exact rescoring +
  sampleFromCandidates. Prompt ingestion also skips the head — prefill
  logits were discarded anyway.
- CLI: enabled automatically on the BitNet path when the head loaded as
  BITNET_PLANES; prints the mode at load.

Measured (real microsoft/BitNet-b1.58-2B-4T I2_S, skainet-cli -s 24,
macOS Apple Silicon, FFM kernels): 6.20-6.29 tok/s vs 4.54-4.72 with the
exact planes matmul and 4.15 with the as-stored F16 tied head — ~+35% over
the planes head, ~+50% over the pre-#357 baseline. Greedy output coherent.

Tests: greedy two-stage decode token-equal to the full-matmul loop
(reference and native stage-1); native stage-1 matches the Kotlin
reference; jvm API dumps regenerated (llm-core's also picks up a stale
generateUntilStop overload entry that predates this change).

Closes #358.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01YKeDSK4JF295y53Uvez954
Base automatically changed from feat/bitnet-tied-lmhead-planes to develop August 31, 2026 17:08
@michalharakal
michalharakal merged commit 2c89e77 into develop Aug 31, 2026
0 of 2 checks passed
@michalharakal
michalharakal deleted the feat/bitnet-two-stage-decode branch August 31, 2026 17:08
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.

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

1 participant