Skip to content

[2/3] feat(dflash): share DFlash masks with inference, default masks for spec_generate - #23

Closed
maocheng23 wants to merge 1 commit into
dflash-is-causal-minfrom
dflash-inference-masks
Closed

maocheng23 wants to merge 1 commit into
dflash-is-causal-minfrom
dflash-inference-masks

Conversation

@maocheng23

@maocheng23 maocheng23 commented Oct 6, 2026 •

Copy link
Copy Markdown
Owner

Split of sgl-project#905, part 2 of 3. Stacked on #22.

Changes

  • Move the mask builders and resolve_dflash_is_causal from algorithms/common/dflash_family_model.py to a new modeling/draft/dflash_mask.py. Names are kept, and dflash_family_model re-imports them, so existing imports, mock patches and benchmark_dflash2_conv.py keep working. Fix DFlash attention masking sgl-project/SpecForge#905 renamed them to build_* and flex_block_size to BLOCK_SIZE; that rename is what broke the benchmark that later landed on main.
  • Qwen3DFlashAttentionBase.is_causal is resolved from the config via resolve_dflash_is_causal(is_causal, config.layer_types[layer_idx]) instead of being hard-coded False.
  • New DFlashDraftModel._prepare_attention_mask: when forward gets no mask (spec_generate), it builds per-layer-type masks for one draft block after cached plus new context.
  • spec_generate no longer forces is_causal=False.

Tests

Ported from sgl-project#905, adapted to the kept names: test_dflash_sliding.py (TestDFlashGenerationMasks), test_dflash_mla.py (spec_generate mask check across eager/sdpa × is_causal × window), and test_dflash_eager_attention.py. The FA2 test now checks that no mask is built, and the reference masks use the left-only in-block window.

CPU run: all tests/test_modeling pass (124 tests, 14 skipped). tests/test_utils gives the same 20 failures as main on this machine (Eagle3 flex tests and 3 modules that can't import here); this PR adds none.

Co-authored with @cih9088 (original sgl-project#905).

🤖 Generated with Claude Code

Move the DFlash mask builders (and resolve_dflash_is_causal) from
algorithms/common/dflash_family_model.py to modeling/draft/dflash_mask.py
so the draft model can use them; dflash_family_model re-imports them, so
existing import paths and test patches keep working.

DFlashDraftModel.forward now builds per-layer-type masks when called
without one (spec_generate): one draft block after the cached plus new
context, following config.is_causal and the sliding window like SGLang.
Bidirectional full layers get no mask. flash_attention_2 gets no mask and
relies on the attention module's is_causal (now resolved from the config
instead of hard-coded False) plus its native window. FA2's window is
symmetric, so for a bidirectional sliding layer with sliding_window <
block_size it also bounds the right side, unlike the explicit masks.
spec_generate no longer forces is_causal=False.

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
maocheng23 force-pushed the dflash-inference-masks branch from 9c404db to e0db52f Compare October 6, 2026 02:48
@maocheng23
maocheng23 force-pushed the dflash-is-causal-min branch from b99ebd4 to 4ba5c24 Compare October 6, 2026 02:48
@maocheng23

Copy link
Copy Markdown
Owner Author

Superseded by sgl-project#935.

@maocheng23 maocheng23 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.

1 participant