From a8c526d419822eb3887f37c00010f4ff9086905c Mon Sep 17 00:00:00 2001 From: maocheng23 Date: Mon, 5 Oct 2026 17:21:19 -0700 Subject: [PATCH] perf(dflash): build flex block masks and run flex attention via compiled 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/SpecForge#905. Co-authored-by: Inhyuk Cho Co-Authored-By: Claude Opus 5.5 --- specforge/modeling/draft/dflash.py | 2 +- specforge/modeling/draft/dflash_mask.py | 11 +++++------ specforge/modeling/draft/flex_attention.py | 2 ++ 3 files changed, 8 insertions(+), 7 deletions(-) diff --git a/specforge/modeling/draft/dflash.py b/specforge/modeling/draft/dflash.py index f298857b1..1ad5b8e95 100644 --- a/specforge/modeling/draft/dflash.py +++ b/specforge/modeling/draft/dflash.py @@ -6,7 +6,6 @@ from torch import nn from transformers import DynamicCache from transformers.cache_utils import Cache -from transformers.integrations.flex_attention import compile_friendly_flex_attention from transformers.modeling_outputs import CausalLMOutputWithPast from transformers.models.qwen3.modeling_qwen3 import ( ALL_ATTENTION_FUNCTIONS, @@ -26,6 +25,7 @@ create_dflash_sdpa_mask, resolve_dflash_is_causal, ) +from .flex_attention import compile_friendly_flex_attention from .flex_attention_backend import flex_attention_backend from .registry import register_draft diff --git a/specforge/modeling/draft/dflash_mask.py b/specforge/modeling/draft/dflash_mask.py index 94c70c5b0..a56f2dfa9 100644 --- a/specforge/modeling/draft/dflash_mask.py +++ b/specforge/modeling/draft/dflash_mask.py @@ -4,11 +4,6 @@ import torch -try: - from torch.nn.attention.flex_attention import create_block_mask -except ImportError: - create_block_mask = None - def resolve_dflash_is_causal(is_causal: Optional[bool], layer_type: str) -> bool: """Resolve in-block causality the way SGLang's DFlash model does. @@ -134,10 +129,14 @@ def dflash_mask_mod(b, h, q_idx, kv_idx): Q_LEN = N * block_size KV_LEN = S + N * block_size + # Imported lazily: the compiled wrapper pulls in torch._dynamo, which the + # dense-mask-only callers (and their sys.modules-patching tests) never need. + from .flex_attention import compile_friendly_create_block_mask + kwargs = {} if flex_block_size is not None: kwargs["BLOCK_SIZE"] = flex_block_size - return create_block_mask( + return compile_friendly_create_block_mask( dflash_mask_mod, B=B, H=None, diff --git a/specforge/modeling/draft/flex_attention.py b/specforge/modeling/draft/flex_attention.py index 50ca5f54d..f4b2c36d3 100644 --- a/specforge/modeling/draft/flex_attention.py +++ b/specforge/modeling/draft/flex_attention.py @@ -89,6 +89,7 @@ def compile_friendly_create_block_mask( Q_LEN, KV_LEN, device, + **kwargs, ): create_block_mask_compiled = ( WrappedCreateBlockMask()() @@ -102,6 +103,7 @@ def compile_friendly_create_block_mask( Q_LEN, KV_LEN, device, + **kwargs, )