Reduce direct-token latency with candidate-only projection - #9
Merged
Merged
Conversation
Cache candidate token indices on the LM-head device and select the requested vocabulary rows before evaluation projection. Retain full-vocabulary training, temperature calibration and the existing device/dtype handling. Add FP32/BF16 candidate-readout parity, training-gradient and derived-buffer checks. The independent branch passes 431 CPU tests (37 optional skips), and builds both a source distribution and wheel. The PR description contains the 4B/27B accuracy and latency comparison with explicit deltas.
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.
Summary
Direct-token inference projects the full vocabulary even though the response uses only the candidate labels. Select the requested LM-head rows before projection and cache the candidate token indices on the head's device, reducing the readout to the number of available choices.
This applies to
lm_tokenevaluation. Training continues to produce full-vocabulary logits; evaluation retains temperature calibration and the existing device/dtype handling. The derived index buffer is rebuilt from checkpoint metadata and is not persisted in checkpoint weights.Accuracy and latency
Oursmeans candidate-only projection plus cached candidate indices. Parentheses show the change relative to the named baseline: accuracy uses percentage points (pp), latency uses percent, and negative latency deltas mean faster. Deltas are calculated from unrounded measurements.Adding Ours changed no answers in either execution mode for these two measured configurations. CUDA Graph itself has small accuracy differences from Eager, shown separately above.
Measurement scope
SimpleJev/JevAny-Qwen3.5-4B-Direct-Token-LoRAcheckpoint, with paired baseline/optimized readouts on the same model. Native graphs were capped at 2,048 tokens: 1,459 requests replayed and 36 used Eager fallback per configuration. The table uses evaluation-traversal timing, which includes first encounters with new shapes and audit-logit transfer.Qwen/Qwen3.8-27Bweights, revision1d4bf0f2ff6012fd82039f2fa52739d0dd7c60c0, using zero-shot candidate scoring in a separate benchmark harness. This is not the trained 27B Pointer release. Both Graph configurations shared the captured text backbone, covered inputs through 4,096 tokens, and replayed all 1,495 requests without fallback. All 665 input-length/candidate-count shapes were warmed before timing; configuration order rotated by request, and audit-logit transfer was excluded. The separate fixed-input panel also measured 8 rounds × 10 calls per configuration and input.Validation
main; onlyjevany/model.py,tests/test_decision_modes.pyandtests/test_backbones.pychange.python -m pytest tests -m 'not server' -q -ra --strict-config --strict-markers— 431 passed, 37 skipped, 7 deselected.python -m build --no-isolation— source distribution and wheel both built successfully.Base:
mainat625aed9.