Skip to content

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

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

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
Owner

Split of sgl-project#905, part 3 of 3. Stacked on #23. Needs a GPU throughput A/B before it goes upstream.

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 sgl-project/SpecForge#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 sgl-project/SpecForge#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 [2/3] feat(dflash): share DFlash masks with inference, default masks for spec_generate #23, 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 with @cih9088 (original sgl-project#905).

🤖 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 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-compiled-block-mask branch from c336279 to a8c526d Compare October 6, 2026 02:48
@maocheng23

Copy link
Copy Markdown
Owner Author

Superseded by sgl-project#936.

@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