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, )