Skip to content

fix(dflash): honor config.is_causal in DFlash training masks [1/3] - #934

Open
maocheng23 wants to merge 1 commit into
mainfrom
dflash-is-causal-min
Open

maocheng23 wants to merge 1 commit into
mainfrom
dflash-is-causal-min

Conversation

@maocheng23

Copy link
Copy Markdown
Collaborator

Split of #905, part 1 of 3 (the actual bug fix). Stacked on upstream main 53398a8.

Motivation

SGLang's DFlash model (v0.5.18+, _get_dflash_attention_type) decides causality inside the draft block from config.is_causal:

  • unset: full layers bidirectional, sliding layers causal
  • false: all layers bidirectional
  • true: all layers causal

SpecForge ignores the field. It always trains sliding layers causal and full layers bidirectional, so any draft with is_causal: false and sliding layers trains causal but is served bidirectional.

Changes

  • create_dflash_sdpa_mask / create_dflash_block_mask take is_causal and resolve it with SGLang's rule via resolve_dflash_is_causal(is_causal, layer_type): an explicit value applies to every layer, unset is the per-layer-type default.
  • The online wrapper passes the draft config's is_causal.
  • Sliding layers also get the window's left bound inside the draft block (kv_off >= q_off - (W-1)). There is no right bound, because SGLang backends disagree on one for non-causal windows (FA3 symmetric (W-1, W-1), FlashInfer/Triton left-only). This only differs from before when sliding_window < block_size, which no shipped config has (W = 2048/4096, block 8/16).
  • benchmark_dflash2_conv.py passes the config's is_causal too.
  • Function names and signatures are otherwise unchanged, so no callers need edits.

⚠️ This changes training for both shipped DFlash2 recipes

configs/qwen3.6-27b-dflash2.json and configs/qwen3.8-27b-dflash2.json set is_causal: false and use only sliding layers. After this PR they train bidirectional inside the block (which is what SGLang v0.5.18+ already serves). This needs a decision before merging:

  • (a) keep as-is, so training follows the config and matches serving;
  • (b) remove is_causal from both configs, which keeps today's causal training and makes serving causal too.

Causal vs bidirectional inside the block has never been ablated cleanly. Existing checkpoints can be served to match how they were trained, without retraining, by deleting is_causal from the served config.json.

Tests

tests/test_utils/test_dflash_mask.py: the reference mask is extended with is_causal. New tests cover dense vs reference, flex mask_mod vs dense, and flex block occupancy for window ∈ {None, 1, 3, 9} × is_causal ∈ {None, False, True}, plus an in-block spot check and a check that the training wrapper reads draft_model.config.is_causal.

Run on CPU (torch 2.14, transformers 5.12.1): every tests/test_utils/test_dflash* and tests/test_modeling/test_dflash* file passes when run alone.

Co-authored-by @cih9088 (original #905); every commit carries a Co-authored-by: Inhyuk Cho trailer.

🤖 Generated with Claude Code

SGLang's DFlash model (v0.5.18+) decides in-block causality from
config.is_causal:
  - unset: full layers bidirectional, sliding layers causal
  - False: every layer bidirectional
  - True:  every layer causal

SpecForge ignored the field and always trained sliding layers causal and
full layers bidirectional, so drafts with `is_causal: false` and sliding
layers were trained causal but served bidirectional.

create_dflash_sdpa_mask / create_dflash_block_mask now take is_causal and
resolve it with SGLang's rule, and the online wrapper passes the draft
config's value. Sliding layers also get the window's left bound inside
the draft block (no right bound: SGLang backends disagree on one for
non-causal windows); this only differs from before when
sliding_window < block_size.

Split out of #905.

Co-authored-by: Inhyuk Cho <ihcho@lgresearch.ai>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

This branch has not been deployed

No deployments
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.

2 participants