feat(bitnet): wire the two-stage lm_head decode into generation (#358) - #362
Merged
Merged
Conversation
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
This was referenced Sep 1, 2026
Closed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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
BitNetTwoStageDecodeandsampleFromCandidatesshipped 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 andpositionadvance exactly asforward(). DIRECT/HYBRID only; OPTIMIZED compiles the full graph and has no partial execution to offer.BitNetStage1Kernel(fun interface) + optional native stage-1 inBitNetTwoStageDecode— one fusedlmhead_stage1call 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 importingTernaryLmheadNative, becauseskainet-backend-apidoes 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 runsforwardHidden: prefill logits were discarded anyway, so every prompt position saves its full-vocab projection outright.BITNET_PLANES; prints "Two-stage lm_head decode: enabled" at load.Measured (real
microsoft/BitNet-b1.58-2B-4TI2_S,skainet-cli -s 24, macOS Apple Silicon, FFM kernels)BITNET_PLANEShead, exact matmul (#361)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
llm-core's also picks up a stalegenerateUntilStopoverload entry that predates this change; the legacyapi/androidfile has no dump task in this configuration and is left as-is).🤖 Generated with Claude Code
https://claude.ai/code/session_01YKeDSK4JF295y53Uvez954