Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion specforge/modeling/draft/dflash.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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

Expand Down
11 changes: 5 additions & 6 deletions specforge/modeling/draft/dflash_mask.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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,
Expand Down
2 changes: 2 additions & 0 deletions specforge/modeling/draft/flex_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,7 @@ def compile_friendly_create_block_mask(
Q_LEN,
KV_LEN,
device,
**kwargs,
):
create_block_mask_compiled = (
WrappedCreateBlockMask()()
Expand All @@ -102,6 +103,7 @@ def compile_friendly_create_block_mask(
Q_LEN,
KV_LEN,
device,
**kwargs,
)


Expand Down
Loading