Skip to content

Fix DFlash attention masking - #905

Closed
cih9088 wants to merge 7 commits into
sgl-project:mainfrom
cih9088:fix/dflash-swa
Closed

cih9088 wants to merge 7 commits into
sgl-project:mainfrom
cih9088:fix/dflash-swa

Conversation

@cih9088

@cih9088 cih9088 commented Sep 28, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

DFlash attention implementation in sglang uses is_causal attribute to determine causal or bidirectional.
https://github.com/sgl-project/sglang/blob/9d18277c2720f3a1d6e64f90259e1d1b9ba9a26a/python/sglang/srt/models/dflash.py#L46-L81

  • is_causal not present: bidirectional for full-attention, causal for sliding-attention
  • is_causal=False: bidirectional for full-attention, sliding-attention
  • is_causal=True: causal for full-attention, sliding-attention

However, SpecForge just ignores the is_causal attribute and always uses bidirectional for full-attention and causal for sliding-attention.

Modifications

  • fixed to follow the sglang attention mask
    • is_causal not present: bidirectional for full-attention, causal for sliding-attention as before
    • is_causal=False: bidirectional for full-attention, sliding-attention
    • is_causal=True: causal for full-attention, sliding-attention
  • used compile friendly version of flex attention and block mask builder from specforge
  • added default attention mask builder method _prepare_attention_mask() to be used for inference (spec_generate())
  • factored the dflash masking builder out to dflash_mask.py so that the masking functions can be used in dflash_model_family.py(training) and dflash.py (inference)

Related Issues

Accuracy Test

Benchmark & Profiling

Checklist

@cih9088 cih9088 changed the title Fix/dflash swa Fix DFlash attention masking Sep 28, 2026
maocheng23 added a commit to maocheng23/SpecForge that referenced this pull request Oct 6, 2026
…led wrappers

create_dflash_block_mask now goes through SpecForge's
compile_friendly_create_block_mask (which forwards BLOCK_SIZE via
**kwargs), and DFlash attention uses SpecForge's
compile_friendly_flex_attention instead of the transformers copy.

The compiled wrapper is imported lazily inside create_dflash_block_mask:
importing it at module scope pulls in torch._dynamo, and tests that build
stubs under patch.dict(sys.modules) then purge it and hit a duplicate
TORCH_LIBRARY registration on re-import (test_dflash_losses run alone).

Note: importing specforge.modeling.draft.flex_attention sets
torch._dynamo.config.recompile_limit = 64 process-wide.

Split out of sgl-project#905.

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

Copy link
Copy Markdown
Collaborator

Thanks @cih9088 for finding and fixing this. The is_causal train/serve mismatch is real. To make review and merge easier, we've split this PR into three stacked PRs, each with you as co-author (Co-authored-by: Inhyuk Cho <ihcho@lgresearch.ai> on every commit):

  1. fix(dflash): honor config.is_causal in DFlash training masks [1/3] #934 [1/3] fix: the core fix. create_dflash_sdpa_mask / create_dflash_block_mask honor config.is_causal with SGLang's rule (resolve_dflash_is_causal(is_causal, layer_type)), and the training wrapper passes it in. Function names are unchanged. The in-block sliding window keeps only the left bound, since SGLang backends disagree on a right bound for non-causal windows (FA3 symmetric, FlashInfer/Triton left-only).
  2. feat(dflash): share DFlash masks with inference, default masks for spec_generate [2/3] #935 [2/3]: moves the builders to modeling/draft/dflash_mask.py (names kept, so benchmark_dflash2_conv.py, which landed on main after this PR, keeps working) and adds the default masks for spec_generate. flash_attention_2 relies on its native causal flag and window instead of raising.
  3. perf(dflash): compiled create_block_mask / flex_attention wrappers [3/3] #936 [3/3, draft]: the compiled create_block_mask / flex_attention wrappers. The import is now lazy, which fixes tests/test_utils/test_dflash_losses.py failing 28/28 when run on its own. It stays a draft until we have a GPU throughput A/B, because mask shapes change every step and recompile.

One open question on #934: configs/qwen3.6-27b-dflash2.json and configs/qwen3.8-27b-dflash2.json set is_causal: false with only sliding layers, so they will start training bidirectional inside the block, which matches SGLang v0.5.18+ serving. Details are in the #934 description.

We'd suggest closing this PR in favor of the split once you've had a look. Feedback on the new PRs is very welcome.

@cih9088 cih9088 closed this Oct 6, 2026
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