Skip to content

perf(dflash): compiled create_block_mask / flex_attention wrappers [3/3] - #936

Open
maocheng23 wants to merge 1 commit into
dflash-inference-masksfrom
dflash-compiled-block-mask
Open

maocheng23 wants to merge 1 commit into
dflash-inference-masksfrom
dflash-compiled-block-mask

Conversation

@maocheng23

@maocheng23 maocheng23 commented Oct 6, 2026 •

Copy link
Copy Markdown
Collaborator

Split of #905, part 3 of 3. Stacked on #935. Draft until a GPU throughput A/B is in.

Changes

  • create_dflash_block_mask uses SpecForge's compile_friendly_create_block_mask (now forwarding **kwargs, so BLOCK_SIZE still works).
  • DFlash attention uses SpecForge's compile_friendly_flex_attention instead of the transformers copy. On torch ≥ 2.7 both are plain torch.compile(flex_attention).
  • The compiled mask wrapper is imported lazily inside create_dflash_block_mask. In Fix DFlash attention masking #905, the module-level import pulled torch._dynamo in under test_dflash_losses's patch.dict(sys.modules). That made the file fail 28/28 when run alone (reproduced on Fix DFlash attention masking #905's head 7ccc036), which CI's single discover run hides. With this PR it passes 28/28 alone.

Open questions

  • Recompiles: DFlash mask shapes (S, N) change every step. On CPU, test_dflash_mask.py goes from 0.1s to 17s because each new shape compiles. GPU training needs a samples/s comparison against feat(dflash): share DFlash masks with inference, default masks for spec_generate [2/3] #935, for example a 200-step Qwen3.8 DFlash2 flex run.
  • Global side effect: importing specforge.modeling.draft.flex_attention sets torch._dynamo.config.recompile_limit = 64 for the whole process. It is now imported by every DFlash user, not only Eagle3/PEagle.

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

🤖 Generated with Claude Code

…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 #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.

1 participant