feat(speculative): support sequence packing in the DSpark Qwen3 draft - #3005
Merged
HuiyingLi merged 11 commits intoJul 12, 2026
Merged
Conversation
DFlash builds its dataloader with build_eagle3_dataloader but never passed
packed_sequence_size, so packing was unreachable. Wire it end to end. DFlash's
Q layout is draft blocks only (context is KV), so packing composes cleanly:
- dflash_mask: create_dflash_{sdpa,block}_mask take optional ctx_doc_id /
anchor_doc_id; a block's context prefix is restricted to its anchor's document
(both the dense SDPA mask and the flex BlockMask mask_mod).
- core: forward accepts position_ids / seq_lens / doc_remaining. Anchor sampling
keeps whole blocks inside one document (doc_remaining >= block_size-1); the
context position ids reset per document; the draft block positions are gathered
from the per-document context positions so RoPE stays document-local.
- target: generate_batch runs the frozen target with the block-causal mask (or FA
varlen at batch size 1) so the captured context hidden states do not leak across
documents, and carries the packing metadata to the trainer.
- train_dflash: read packed_sequence_size into both dataloaders, splat the packing
metadata into generate_batch and the trainer, fail fast on incompatible configs
(cp_size>1, FlashAttention target with batch>1), and reject packing on the
JetSpec/Domino subclasses (they override the trainer forward without threading it).
Tests: doc-id construction, SDPA mask document restriction, anchor sampling stays
in-document, per-document draft positions, single-document packed==unpacked
equivalence, a two-document packed training step, target-side document isolation,
and the recipe helpers/gates.
Signed-off-by: khazic <khazzz1c@gmail.com>
Signed-off-by: khazic <khazzz1c@gmail.com>
… limitation Address /code-review findings: the flex_attention backend is the DFlash default but every packing test used sdpa, leaving the flex doc-gating branch uncovered. Factor the flex mask_mod into a module-level build_dflash_mask_mod so it can be materialised on CPU via create_mask (FlexAttention's kernel is CUDA-only), and add a test asserting the flex per-token mask equals the SDPA mask and enforces the per-document restriction. Also document that the whole-block-in-document anchor bound (required because _build_block_targets does not encode boundaries) drops packed documents shorter than block_size. Signed-off-by: khazic <khazzz1c@gmail.com>
Signed-off-by: khazic <khazzz1c@gmail.com>
DSpark is anchor/block-based like DFlash and reuses the same dflash masks and the shared dspark/common.py helpers, so packing composes the same way. Scope this to the online text-only Qwen3 path: - common.py: sample_anchor_positions / build_anchor_candidate_mask require the anchor's first target to stay in its document; build_eval_mask truncates each block at its anchor's document boundary (a partial in-document block, matching the sequence-end truncation); create_position_ids gathers per-document context positions; add a context_doc_ids helper. All new params default to None so the other drafts' calls are unchanged. - draft_qwen3.py: forward accepts position_ids / seq_lens / doc_remaining, builds the document-aware dflash mask (ctx_doc_id/anchor_doc_id) and per-document positions, and gates supervision by document. - core.py / target.py: thread the packing metadata; the target runs block-causal (or FA varlen at batch size 1) so context hidden states do not leak across docs. - train_dspark.py: read packed_sequence_size into the LLM dataloaders, splat the metadata through _forward_batch, and fail fast on unsupported configs (multimodal, cached-target, non-Qwen3 draft, cp_size>1, FA target with batch>1). The Markov head is position-agnostic (token-id lookup) and needs no change. Tests cover the doc-aware common helpers, single-document packed==unpacked equivalence, a two-document training step, target-side document isolation, and the recipe gates. Signed-off-by: khazic <khazzz1c@gmail.com>
Signed-off-by: khazic <khazzz1c@gmail.com>
… guard target position_ids Address /code-review findings on the DSpark packing change: - DSparkTrainerModule.forward passed the new position_ids/seq_lens/doc_remaining kwargs to the draft unconditionally, but only the Qwen3 draft accepts them, so every non-Qwen3 DSpark draft (DeepseekV4/Gemma4/GLM-5.2/MiniMax-M3) raised a TypeError on the first training step even on the default unpacked path. Pass the packing metadata only when packing is active (seq_lens is not None); packing is gated to the Qwen3 draft at recipe setup. - The target wrapper builds its forward kwargs through filter_forward_kwargs, which silently drops kwargs the target does not declare. Under packing, fail loud if the target forward does not accept position_ids (dropping it would give wrong per-document RoPE and, on the FlashAttention path, lose the document-boundary signal) instead of silently leaking context across documents. Tests: the unpacked trainer must not forward packing kwargs to a legacy-signature draft. Signed-off-by: khazic <khazzz1c@gmail.com>
khazic
force-pushed
the
khazic/feat/dspark-draft-packing
branch
from
July 10, 2026 18:26
9e6e840 to
3e6c934
Compare
This was referenced Jul 10, 2026
Signed-off-by: khazic <khazzz1c@gmail.com> # Conflicts: # nemo_automodel/components/speculative/dflash/core.py
…draft-packing # Conflicts: # nemo_automodel/components/speculative/dflash/target.py
HuiyingLi
approved these changes
Jul 12, 2026
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.
What
Extends sequence packing to the DSpark draft. DSpark is anchor/block-based like DFlash and reuses the same dflash masks plus the shared
dspark/common.pyhelpers, so packing composes the same way. Scope: the online text-only Qwen3 path.sample_anchor_positions/build_anchor_candidate_maskrequire the anchor's first target to stay in its document;build_eval_masktruncates each block at its anchor's document boundary (step <= doc_remaining[anchor]before the cumprod, a partial in-document block matching the sequence-end truncation);create_position_idsgathers per-document context positions; add acontext_doc_idshelper. All new params default toNone, so the other drafts' calls are unchanged.forwardacceptsposition_ids/seq_lens/doc_remaining, builds the document-aware dflash mask (ctx_doc_id/anchor_doc_id) and per-document positions, and gates supervision by document.position_ids.packed_sequence_sizeinto the LLM dataloaders, splat the metadata through_forward_batch, and fail fast on unsupported configs (multimodal, cached-target, non-Qwen3 draft,cp_size>1, FlashAttention target withmicro_batch_size>1).The Markov head is position-agnostic (token-id lookup) and needs no change. The other 4 DSpark drafts + the VLM and cached paths are gated off (documented follow-ups).
Correctness verification
CPU unit tests (14 pass; the packing helpers run on the SDPA/eager path), including:
seq_lens; anchor candidate requires the first target in-document; eval_mask truncates a block at the document boundary; per-document draft positions;GPU recipe smoke (1x A100-80GB, bf16), the full DSpark recipe with a real Qwen3-4B target, packing enabled, and the flex_attention backend:
Training log
The three-term loss falls with packing on, exercising the full path (packed dataloader -> target block-causal -> document-aware flex block mask -> anchor / position / eval-mask handling -> Markov head).
Scope
DSpark Qwen3 online text-only path only. The other 4 DSpark drafts, the VLM path, the offline-cache path, and packing with context parallelism remain follow-ups.