From 4ba5c24bece5b9e8d4d207e8598719f1775a6771 Mon Sep 17 00:00:00 2001 From: maocheng23 Date: Mon, 5 Oct 2026 17:15:25 -0700 Subject: [PATCH] fix(dflash): honor config.is_causal in DFlash training masks SGLang's DFlash model (v0.5.18+) decides in-block causality from config.is_causal: - unset: full layers bidirectional, sliding layers causal - False: every layer bidirectional - True: every layer causal SpecForge ignored the field and always trained sliding layers causal and full layers bidirectional, so drafts with `is_causal: false` and sliding layers were trained causal but served bidirectional. create_dflash_sdpa_mask / create_dflash_block_mask now take is_causal and resolve it with SGLang's rule, and the online wrapper passes the draft config's value. Sliding layers also get the window's left bound inside the draft block (no right bound: SGLang backends disagree on one for non-causal windows); this only differs from before when sliding_window < block_size. Split out of sgl-project/SpecForge#905. Co-authored-by: Inhyuk Cho Co-Authored-By: Claude Opus 5.5 --- .../algorithms/common/dflash_family_model.py | 44 +++++- .../benchmarks/benchmark_dflash2_conv.py | 1 + tests/test_utils/test_dflash_mask.py | 128 +++++++++++++++++- 3 files changed, 164 insertions(+), 9 deletions(-) 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]])