diff --git a/specforge/algorithms/common/dflash_family_model.py b/specforge/algorithms/common/dflash_family_model.py index 719b12232..98b76fc78 100644 --- a/specforge/algorithms/common/dflash_family_model.py +++ b/specforge/algorithms/common/dflash_family_model.py @@ -245,6 +245,18 @@ 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, @@ -252,11 +264,16 @@ def create_dflash_sdpa_mask( 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 @@ -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) @@ -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 @@ -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 @@ -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" diff --git a/specforge/benchmarks/benchmark_dflash2_conv.py b/specforge/benchmarks/benchmark_dflash2_conv.py index f41d0f52e..bd858a739 100644 --- a/specforge/benchmarks/benchmark_dflash2_conv.py +++ b/specforge/benchmarks/benchmark_dflash2_conv.py @@ -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) diff --git a/tests/test_utils/test_dflash_mask.py b/tests/test_utils/test_dflash_mask.py index d4f71f069..17b7e8331 100644 --- a/tests/test_utils/test_dflash_mask.py +++ b/tests/test_utils/test_dflash_mask.py @@ -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 @@ -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 @@ -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) @@ -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( @@ -110,6 +119,7 @@ def _compare_masks( block_size=block_size, device=self.device, sliding_window=sliding_window, + is_causal=is_causal, ) self.assertEqual( @@ -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()}", ) @@ -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) @@ -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( @@ -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) @@ -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( @@ -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]])