Skip to content

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

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

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

Conversation

@maocheng23

Copy link
Copy Markdown
Collaborator

Split of #905, part 2 of 3. Stacked on #934.

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 #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.
    • Bidirectional full layers get no mask.
    • eager / sdpa / flex get explicit masks.
    • flash_attention_2 gets no mask and relies on module.is_causal plus FA's native window with bottom-right causal alignment. Fix DFlash attention masking #905 raised NotImplementedError for FA2 here. FA2's window is symmetric (W-1, W-1), so for a bidirectional sliding layer with sliding_window < block_size it also bounds the right side, unlike the explicit masks (left bound only, see fix(dflash): honor config.is_causal in DFlash training masks [1/3] #934). This is an edge case only spec_generate can reach.
  • spec_generate no longer forces is_causal=False.

Tests

Ported from #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-by @cih9088 (original #905); every commit carries a Co-authored-by: Inhyuk Cho trailer.

🤖 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 #905.

Co-authored-by: Inhyuk Cho <ihcho@lgresearch.ai>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
mask_builder = create_dflash_block_mask
if flex_attention_backend() == "FLASH":
# FLASH requires a minimum of this block size.
mask_args["flex_block_size"] = (256, 128)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
mask_args["flex_block_size"] = (256, 128)
mask_args["BLOCK_SIZE"] = (256, 128)

how about just passing additional kwargs to the block mask creation function directly, instead of introducing merely a wrapper argument?

Comment on lines +137 to +139
kwargs = {}
if flex_block_size is not None:
kwargs["BLOCK_SIZE"] = flex_block_size

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
kwargs = {}
if flex_block_size is not None:
kwargs["BLOCK_SIZE"] = flex_block_size

how about just passing additional kwargs to the block mask creation function directly, instead of introducing merely a wrapper argument?

Comment on lines +90 to +92
flex_block_size=None,
sliding_window: Optional[int] = None,
is_causal: Optional[bool] = None,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
flex_block_size=None,
sliding_window: Optional[int] = None,
is_causal: Optional[bool] = None,
sliding_window: Optional[int] = None,
is_causal: Optional[bool] = None,
**kwargs,

how about just passing additional kwargs to the block mask creation function directly, instead of introducing merely a wrapper argument?

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