Repository navigation
perf(qwen3-omni): project only final prefill logits - #446
Merged
CjhHa1 merged 5 commits intoSep 17, 2026
Merged
Conversation
CjhHa1
requested review from
celve,
haonan3 and
zzhuoxin1508
as code owners
September 13, 2026 04:19
CjhHa1
force-pushed
the
perf/selective-trainside-logits
branch
from
September 13, 2026 04:22
f4f275d to
2bd764c
Compare
CjhHa1
force-pushed
the
perf/selective-trainside-logits
branch
from
September 14, 2026 01:13
2bd764c to
825b6da
Compare
1 task done
celve
approved these changes
Sep 16, 2026
2 tasks
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
Avoid materializing prompt-wide vocabulary logits in the Qwen3-Omni trainside prefill when generation samples only the final row.
The current Transformers 5.6 Qwen3-Omni Thinker forward does not expose
logits_to_keep. A scoped helper selectshidden_states[:, -1:, :]at thelm_headboundary without replacing the module or changing its state dict. The hook exists only around the first prefill forward; decode and replay remain unchanged.Related Issue
N/A
Test Plan
SKIP=no-commit-to-branch pre-commit run --all-files --show-diff-on-failure— pass[B, 1, V], matches the full projection's final row, and removes its hook after the forwardB=4,L=512,H=2048,V=152064, bf16, taskselective-logits-head-bench-0913): peak allocation 0.580 GiB -> 1.160 MiB; median 9.09 ms -> 0.202 ms; selected logits exact[4, 512](taskpr446-full-prefill-ab-8h20-0913-v3): max per-GPU forward allocation 0.596 -> 0.307 GiB (-0.289 GiB, -48%); median 259.14 -> 254.48 ms (+1.8%); selected logits close and sampled KV cache exact[4, 512](taskpr446-fsdp-qwen-8h20-0914): max-rank forward peak 3.958 -> 3.958 GiB (unchanged because FSDP all-gather dominates); median max-rank latency 363.03 -> 344.13 ms (+5.5%); selected logits close and sampled KV cache exact on every rankCompatibility / Risk
No config, checkpoint, state-dict, or public model API changes. Prompts are left-padded before generation, so the final hidden row is each sample's last real token. The forward pre-hook is removed in
finally; subsequent decode steps and replay do not carry it.Under production FSDP the optimization improves prefill latency, but does not reduce the observed peak because the larger FSDP layer all-gather determines that peak.
Reviewer Notes
The native Qwen
logits_to_keepchanges live in #442. This PR covers Qwen3-Omni because its current forward lacks that argument.HunyuanImage3 was benchmarked and removed from this PR: under its production FSDP topology, peak changed only 9.474 -> 9.471 GiB and latency regressed 2672.10 -> 2707.63 ms, so the scoped projection did not justify its added complexity.
Checklist