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
44 changes: 40 additions & 4 deletions specforge/algorithms/common/dflash_family_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -245,18 +245,35 @@ def compute_walk_accepted_length_terms(
return accepted.new_tensor((accepted_total, visited_count)).unbind()


def resolve_dflash_is_causal(is_causal: Optional[bool], layer_type: str) -> bool:
"""Resolve in-block causality the way SGLang's DFlash model does.

An explicit ``is_causal`` applies to every layer. Unset is a per-layer-type
default: full layers bidirectional, sliding layers causal.
"""

if is_causal is not None:
return bool(is_causal)
return layer_type == "sliding_attention"


def create_dflash_sdpa_mask(
anchor_positions,
block_keep_mask,
S,
block_size,
device,
sliding_window: Optional[int] = None,
is_causal: Optional[bool] = None,
):
"""Construct a full or sliding dense boolean DFlash mask."""

if sliding_window is not None and sliding_window <= 0:
raise ValueError("sliding_window must be > 0")
is_causal = resolve_dflash_is_causal(
is_causal,
"sliding_attention" if sliding_window is not None else "full_attention",
)
B, N = anchor_positions.shape
Q_LEN = N * block_size
KV_LEN = S + N * block_size
Expand All @@ -282,9 +299,16 @@ def create_dflash_sdpa_mask(
is_draft = kv_indices >= S
kv_block_ids = (kv_indices - S) // block_size
mask_draft = is_draft & (q_block_ids == kv_block_ids)
if sliding_window is not None:
kv_block_offsets = (kv_indices - S) % block_size
kv_block_offsets = (kv_indices - S) % block_size
if is_causal:
mask_draft = mask_draft & (kv_block_offsets <= q_block_offsets)
if sliding_window is not None:
# Left window bound inside the block. Only bites when sliding_window <
# block_size; no right bound, since SGLang backends disagree on one for
# non-causal windows (FA3 symmetric, FlashInfer/Triton left-only).
mask_draft = mask_draft & (
kv_block_offsets >= q_block_offsets - (sliding_window - 1)
)

valid_block = block_keep_mask.view(B, 1, N, 1).repeat_interleave(block_size, dim=2)

Expand All @@ -300,11 +324,16 @@ def create_dflash_block_mask(
device: torch.device,
flex_block_size=None,
sliding_window: Optional[int] = None,
is_causal: Optional[bool] = None,
):
"""Construct a full or sliding Flex Attention mask for DFlash training."""

if sliding_window is not None and sliding_window <= 0:
raise ValueError("sliding_window must be > 0")
is_causal = resolve_dflash_is_causal(
is_causal,
"sliding_attention" if sliding_window is not None else "full_attention",
)

def dflash_mask_mod(b, h, q_idx, kv_idx):
q_block_id = q_idx // block_size
Expand All @@ -324,9 +353,13 @@ def dflash_mask_mod(b, h, q_idx, kv_idx):
is_draft = kv_idx >= S
kv_block_id = (kv_idx - S) // block_size
mask_draft = is_draft & (q_block_id == kv_block_id)
if sliding_window is not None:
kv_block_offset = (kv_idx - S) % block_size
kv_block_offset = (kv_idx - S) % block_size
if is_causal:
mask_draft = mask_draft & (kv_block_offset <= q_block_offset)
if sliding_window is not None:
mask_draft = mask_draft & (
kv_block_offset >= q_block_offset - (sliding_window - 1)
)

is_valid_block = block_keep_mask[b, safe_q_block_id]
in_bounds = q_block_id < N
Expand Down Expand Up @@ -732,6 +765,9 @@ def _forward_draft_blocks(
"S": seq_len,
"block_size": self.block_size,
"device": device,
"is_causal": getattr(
getattr(self.draft_model, "config", None), "is_causal", None
),
}
if (
self.attention_backend == "flex_attention"
Expand Down
1 change: 1 addition & 0 deletions specforge/benchmarks/benchmark_dflash2_conv.py
Original file line number Diff line number Diff line change
Expand Up @@ -271,6 +271,7 @@ def benchmark_draft(model, config, args):
S=seq_len,
block_size=block_size,
device=device,
is_causal=getattr(model.config, "is_causal", None),
)
if flex_attention_backend() == "FLASH":
mask_args["flex_block_size"] = (256, 128)
Expand Down
128 changes: 123 additions & 5 deletions tests/test_utils/test_dflash_mask.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,11 +22,15 @@ def _reference_dflash_mask(
block_size,
device,
sliding_window=None,
is_causal=None,
):
"""Element-level reference for full and sliding DFlash attention.

This uses plain Python loops so correctness is obvious by inspection.
"""
if is_causal is None:
# SGLang: unset -> only sliding layers are causal inside the block.
is_causal = sliding_window is not None
B, N = anchor_positions.shape
Q_LEN = N * block_size
KV_LEN = S + N * block_size
Expand All @@ -47,14 +51,17 @@ def _reference_dflash_mask(
is_draft = kv_idx >= S
kv_block_id = (kv_idx - S) // block_size
draft_visible = is_draft and (q_block_id == kv_block_id)
kv_offset = (kv_idx - S) % block_size
if is_causal:
draft_visible = draft_visible and kv_offset <= q_offset

if sliding_window is not None:
q_offset = q_idx % block_size
kv_offset = (kv_idx - S) % block_size
ctx_visible = ctx_visible and (
kv_idx >= anchor_pos + q_offset - (sliding_window - 1)
)
draft_visible = draft_visible and kv_offset <= q_offset
draft_visible = draft_visible and (
kv_offset >= q_offset - (sliding_window - 1)
)

if ctx_visible or draft_visible:
mask[b, 0, q_idx, kv_idx] = True
Expand Down Expand Up @@ -89,6 +96,7 @@ def _compare_masks(
S,
block_size,
sliding_window=None,
is_causal=None,
):
"""Compare create_dflash_sdpa_mask against element-level reference (ground truth)."""
anchor_positions = anchor_positions.to(self.device)
Expand All @@ -101,6 +109,7 @@ def _compare_masks(
block_size=block_size,
device=self.device,
sliding_window=sliding_window,
is_causal=is_causal,
)

ref_mask = _reference_dflash_mask(
Expand All @@ -110,6 +119,7 @@ def _compare_masks(
block_size=block_size,
device=self.device,
sliding_window=sliding_window,
is_causal=is_causal,
)

self.assertEqual(
Expand All @@ -120,7 +130,7 @@ def _compare_masks(
self.assertTrue(
torch.equal(sdpa_mask, ref_mask),
f"Mask mismatch with S={S}, block_size={block_size}, "
f"sliding_window={sliding_window}, anchors={anchor_positions.tolist()}, "
f"sliding_window={sliding_window}, is_causal={is_causal}, anchors={anchor_positions.tolist()}, "
f"keep={block_keep_mask.tolist()}\n"
f"Diff positions: {(sdpa_mask != ref_mask).nonzero(as_tuple=False).tolist()}",
)
Expand All @@ -132,6 +142,7 @@ def _compare_block_mask_consistency(
S,
block_size,
sliding_window=None,
is_causal=None,
):
"""Verify create_dflash_block_mask block-level mask is consistent with reference."""
anchor_positions = anchor_positions.to(self.device)
Expand All @@ -144,6 +155,7 @@ def _compare_block_mask_consistency(
block_size=block_size,
device=self.device,
sliding_window=sliding_window,
is_causal=is_causal,
)

ref_mask = _reference_dflash_mask(
Expand All @@ -153,6 +165,7 @@ def _compare_block_mask_consistency(
block_size=block_size,
device=self.device,
sliding_window=sliding_window,
is_causal=is_causal,
)

dense_blocks = block_mask.to_dense() # (B, H, Q_blocks, KV_blocks)
Expand Down Expand Up @@ -257,7 +270,7 @@ def test_sliding_window_moves_with_query_offset(self):
self.assertEqual(actual, expected)

def test_sliding_window_one_has_no_context_and_causal_draft(self):
"""A one-token window removes context but keeps causal own-block keys."""
"""A one-token window removes context and leaves only the own draft slot."""
anchor_positions = torch.tensor([[4, 9]])
block_keep_mask = torch.tensor([[True, True]])
self._compare_masks(
Expand Down Expand Up @@ -324,6 +337,111 @@ def test_all_dflash_families_build_mixed_layer_masks(self):
self.assertTrue(masks["full_attention"][0, 0, 0, 0].item())
self.assertFalse(masks["sliding_attention"][0, 0, 0, 0].item())

def test_is_causal_follows_sglang_rule(self):
"""Unset/False/True ``is_causal`` match SGLang's per-layer-type rule."""
anchors = torch.tensor([[5, 11]])
keep = torch.tensor([[True, True]])
for sliding_window in (None, 1, 3, 9):
for is_causal in (None, False, True):
with self.subTest(sliding_window=sliding_window, is_causal=is_causal):
self._compare_masks(
anchors,
keep,
S=16,
block_size=4,
sliding_window=sliding_window,
is_causal=is_causal,
)
self._compare_block_mask_consistency(
anchors,
keep,
S=16,
block_size=4,
sliding_window=sliding_window,
is_causal=is_causal,
)
mask_args = {
"anchor_positions": anchors.to(self.device),
"block_keep_mask": keep.to(self.device),
"S": 16,
"block_size": 4,
"device": self.device,
"sliding_window": sliding_window,
"is_causal": is_causal,
}
dense_mask = create_dflash_sdpa_mask(**mask_args)
block_mask = create_dflash_block_mask(**mask_args)
flex_mask = block_mask.mask_mod(
0,
0,
torch.arange(8, device=self.device).unsqueeze(1),
torch.arange(24, device=self.device).unsqueeze(0),
)
self.assertTrue(torch.equal(flex_mask, dense_mask[0, 0]))

def test_in_block_visibility_per_is_causal(self):
"""Spot-check the draft block: offset 0 -> 3 (right) and 3 -> 0 (left)."""
anchors = torch.tensor([[6]], device=self.device)
keep = torch.tensor([[True]], device=self.device)
# (sliding_window, is_causal, sees_right, sees_far_left)
cases = (
(None, None, True, True), # full, unset -> bidirectional
(None, False, True, True),
(None, True, False, True),
(8, None, False, True), # sliding, unset -> causal
(8, False, True, True),
(8, True, False, True),
(2, False, True, False), # window bounds the left side only
(2, True, False, False),
)
for sliding_window, is_causal, sees_right, sees_far_left in cases:
with self.subTest(sliding_window=sliding_window, is_causal=is_causal):
mask = create_dflash_sdpa_mask(
anchor_positions=anchors,
block_keep_mask=keep,
S=12,
block_size=4,
device=self.device,
sliding_window=sliding_window,
is_causal=is_causal,
)
self.assertEqual(mask[0, 0, 0, 12 + 3].item(), sees_right)
self.assertEqual(mask[0, 0, 3, 12 + 0].item(), sees_far_left)
self.assertTrue(mask[0, 0, 3, 12 + 3].item())

def test_training_mask_reads_draft_config_is_causal(self):
anchors = torch.tensor([[12]], device=self.device)
keep = torch.tensor([[True]], device=self.device)
for is_causal in (None, False, True):
with self.subTest(is_causal=is_causal):
draft_model = _RecordingDraftModel().to(self.device)
if is_causal is not None:
draft_model.config.is_causal = is_causal
model = OnlineDFlashModel(
draft_model=draft_model,
target_lm_head=nn.Identity(),
target_embed_tokens=nn.Embedding(32, 8).to(self.device),
mask_token_id=31,
block_size=4,
attention_backend="sdpa",
)
with mock.patch.object(
model,
"_sample_anchor_positions",
return_value=(anchors, keep),
):
model._forward_draft_blocks(
input_ids=torch.arange(16, device=self.device).unsqueeze(0),
hidden_states=torch.randn(1, 16, 8, device=self.device),
loss_mask=torch.ones(1, 16, device=self.device),
)
masks = draft_model.attention_mask
# Query offset 0 of the block vs. key offset 1 (kv index S + 1).
full_future = masks["full_attention"][0, 0, 0, 17].item()
sliding_future = masks["sliding_attention"][0, 0, 0, 17].item()
self.assertEqual(full_future, not is_causal)
self.assertEqual(sliding_future, is_causal is False)

def test_mixed_validity_multi_batch(self):
"""Multi-batch with mixed block validity patterns."""
anchor_positions = torch.tensor([[10, 40, 70, 100], [20, 50, 80, 110]])
Expand Down
Loading