From f9b3c182f3185171935b0332bf5524d8a3e4ef05 Mon Sep 17 00:00:00 2001 From: Rahul Chalamala <22563365+rchalamala@users.noreply.github.com> Date: Tue, 22 Sep 2026 16:35:18 -0700 Subject: [PATCH 1/2] dflash: keep grouped-conv taps inside each request's block (#148) * dflash: keep grouped-conv taps inside each request's block _grouped_conv shifts rows across the flattened [bs * block_size] token dimension and masked the taps that cross a block boundary by multiplying with a 0/1 mask. NaN * 0 and Inf * 0 are NaN, so a non-finite value in request i's last block rows reached request i+1's first rows, and the DFlash2 draft stack (attention_conv and mlp_conv, prepare and finish, in every layer) carried it one request further per layer. Select the cross-block taps to exact zeros with torch.where before the multiply. Finite outputs are unchanged apart from the sign of exact zeros. * dflash: test the compiled grouped conv at the block boundary The engine calls _grouped_conv through torch.compile, but every boundary test forced the eager original with set_stance("force_eager"), so CI never ran the compiled function on non-finite input. Run the boundary check and the DFlashGroupedConv prepare/finish check both eagerly and through the compiled function (inductor on the CPU runner). With the multiply-by-mask formulation restored in dflash.py, the compiled checks fail at the same boundary rows as the eager ones (taps 2 row 7, taps 3 rows 6-7, for NaN, +Inf and -Inf) and in prepare/finish. The first CPU compile takes about 20 s, so est_time goes to 30. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- python/sglang/srt/models/dflash.py | 6 +- .../unit/spec/test_dflash_grouped_conv.py | 114 ++++++++++++++++++ 2 files changed, 119 insertions(+), 1 deletion(-) create mode 100644 test/registered/unit/spec/test_dflash_grouped_conv.py diff --git a/python/sglang/srt/models/dflash.py b/python/sglang/srt/models/dflash.py index f3205f3de4ae..767ec06b2a14 100644 --- a/python/sglang/srt/models/dflash.py +++ b/python/sglang/srt/models/dflash.py @@ -472,8 +472,12 @@ def _grouped_conv(hidden_states, delta, base, block_size, num_groups, group_size else: position = position % block_size for tap in range(1, taps): + # Taps that would reach into the previous block (another request) must be + # exact zeros. Multiplying by a 0/1 mask is not enough: NaN or Inf in the + # previous request's last rows would become NaN here (NaN * 0 = NaN). shifted = F.pad(blocks[:-tap], (0, 0, 0, 0, tap, 0)) - out = out + coefficients[:, tap] * shifted * (position >= tap).view(-1, 1, 1) + shifted = torch.where((position >= tap).view(-1, 1, 1), shifted, 0) + out = out + coefficients[:, tap] * shifted return out.flatten(-2) diff --git a/test/registered/unit/spec/test_dflash_grouped_conv.py b/test/registered/unit/spec/test_dflash_grouped_conv.py new file mode 100644 index 000000000000..834e26002cba --- /dev/null +++ b/test/registered/unit/spec/test_dflash_grouped_conv.py @@ -0,0 +1,114 @@ +"""DFlash grouped conv: a block's taps must never read another request's rows.""" + +import contextlib +import unittest + +import torch +import torch.nn.functional as F + +from sglang.srt.models.dflash import DFlashGroupedConv, _grouped_conv +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=30, suite="base-a-test-cpu") + +HIDDEN, GROUP, BLOCK = 64, 16, 8 +NUM_GROUPS = HIDDEN // GROUP + + +def _inputs(bs, taps, dtype, seed=0): + g = torch.Generator().manual_seed(seed) + x = torch.randn(bs * BLOCK, HIDDEN, generator=g).to(dtype) + delta = (0.5 * torch.randn(bs * BLOCK, taps, NUM_GROUPS, generator=g)).to(dtype) + base = torch.randn(taps, HIDDEN, generator=g).to(dtype) + return x, delta, base + + +def _stance(compiled): + # The engine calls the torch.compile'd _grouped_conv (inductor here); the + # eager stance runs its original Python instead. + if compiled: + return contextlib.nullcontext() + return torch.compiler.set_stance("force_eager") + + +def _conv(x, delta, base, taps, compiled=False): + with _stance(compiled): + return _grouped_conv(x, delta, base, BLOCK, NUM_GROUPS, GROUP, taps) + + +def _multiplicative_mask_conv(x, delta, base, taps): + # The previous formulation, kept to show finite outputs are unchanged. + blocks = x.unflatten(-1, (NUM_GROUPS, GROUP)) + coefficients = base.view(1, taps, NUM_GROUPS, GROUP) + delta.unsqueeze(-1) + out = coefficients[:, 0] * blocks + position = torch.arange(x.shape[0]) % BLOCK + for tap in range(1, taps): + shifted = F.pad(blocks[:-tap], (0, 0, 0, 0, tap, 0)) + out = out + coefficients[:, tap] * shifted * (position >= tap).view(-1, 1, 1) + return out.flatten(-2) + + +class TestDFlashGroupedConv(CustomTestCase): + def _check_nonfinite_rows_stay_in_their_request(self, compiled): + for taps in (2, 3): + for bad in (float("nan"), float("inf"), float("-inf")): + for row in range(BLOCK): + with self.subTest(taps=taps, bad=bad, row=row): + x, delta, base = _inputs(bs=3, taps=taps, dtype=torch.bfloat16) + clean = _conv(x, delta, base, taps, compiled) + clean = clean.view(3, BLOCK, HIDDEN) + x[row, 5] = bad # request 0 + out = _conv(x, delta, base, taps, compiled) + out = out.view(3, BLOCK, HIDDEN) + torch.testing.assert_close(out[1:], clean[1:], rtol=0, atol=0) + bad_rows = (~torch.isfinite(out[0])).any(-1).nonzero() + expected = torch.arange(row, min(row + taps, BLOCK)) + self.assertEqual(bad_rows.flatten().tolist(), expected.tolist()) + + def _check_module_prepare_and_finish_do_not_leak(self, compiled): + torch.manual_seed(0) + conv = DFlashGroupedConv(HIDDEN, BLOCK, 2, GROUP).to(torch.bfloat16) + x = torch.randn(2 * BLOCK, HIDDEN).to(torch.bfloat16) + y = torch.randn(2 * BLOCK, HIDDEN).to(torch.bfloat16) + with torch.no_grad(), _stance(compiled): + ref_x, kernel = conv.prepare(x) + ref_y = conv.finish(y, kernel) + x_bad, y_bad = x.clone(), y.clone() + x_bad[BLOCK - 1] = float("nan") + y_bad[BLOCK - 1] = float("inf") + out_x, _ = conv.prepare(x_bad) + out_y = conv.finish(y_bad, kernel) + torch.testing.assert_close(out_x[BLOCK:], ref_x[BLOCK:], rtol=0, atol=0) + torch.testing.assert_close(out_y[BLOCK:], ref_y[BLOCK:], rtol=0, atol=0) + + def test_nonfinite_rows_stay_in_their_request(self): + self._check_nonfinite_rows_stay_in_their_request(compiled=False) + + def test_compiled_nonfinite_rows_stay_in_their_request(self): + self._check_nonfinite_rows_stay_in_their_request(compiled=True) + + def test_module_prepare_and_finish_do_not_leak(self): + self._check_module_prepare_and_finish_do_not_leak(compiled=False) + + def test_compiled_module_prepare_and_finish_do_not_leak(self): + self._check_module_prepare_and_finish_do_not_leak(compiled=True) + + def test_finite_output_unchanged(self): + for dtype in (torch.float32, torch.bfloat16): + for taps in (2, 3): + for seed in range(4): + with self.subTest(dtype=dtype, taps=taps, seed=seed): + x, delta, base = _inputs(5, taps, dtype, seed) + # torch.equal treats -0.0 == 0.0: the old mask could + # produce -0.0 where the new one produces +0.0. + self.assertTrue( + torch.equal( + _conv(x, delta, base, taps), + _multiplicative_mask_conv(x, delta, base, taps), + ) + ) + + +if __name__ == "__main__": + unittest.main() From db2f8b8bb1e7dadcc767a476ef7e0068e3127f62 Mon Sep 17 00:00:00 2001 From: rdamani Date: Thu, 1 Oct 2026 02:35:03 +0000 Subject: [PATCH 2/2] dflash: drop grouped-conv unit test and inline comment per review Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- python/sglang/srt/models/dflash.py | 3 - .../unit/spec/test_dflash_grouped_conv.py | 114 ------------------ 2 files changed, 117 deletions(-) delete mode 100644 test/registered/unit/spec/test_dflash_grouped_conv.py diff --git a/python/sglang/srt/models/dflash.py b/python/sglang/srt/models/dflash.py index 767ec06b2a14..05b6f1f8d857 100644 --- a/python/sglang/srt/models/dflash.py +++ b/python/sglang/srt/models/dflash.py @@ -472,9 +472,6 @@ def _grouped_conv(hidden_states, delta, base, block_size, num_groups, group_size else: position = position % block_size for tap in range(1, taps): - # Taps that would reach into the previous block (another request) must be - # exact zeros. Multiplying by a 0/1 mask is not enough: NaN or Inf in the - # previous request's last rows would become NaN here (NaN * 0 = NaN). shifted = F.pad(blocks[:-tap], (0, 0, 0, 0, tap, 0)) shifted = torch.where((position >= tap).view(-1, 1, 1), shifted, 0) out = out + coefficients[:, tap] * shifted diff --git a/test/registered/unit/spec/test_dflash_grouped_conv.py b/test/registered/unit/spec/test_dflash_grouped_conv.py deleted file mode 100644 index 834e26002cba..000000000000 --- a/test/registered/unit/spec/test_dflash_grouped_conv.py +++ /dev/null @@ -1,114 +0,0 @@ -"""DFlash grouped conv: a block's taps must never read another request's rows.""" - -import contextlib -import unittest - -import torch -import torch.nn.functional as F - -from sglang.srt.models.dflash import DFlashGroupedConv, _grouped_conv -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=30, suite="base-a-test-cpu") - -HIDDEN, GROUP, BLOCK = 64, 16, 8 -NUM_GROUPS = HIDDEN // GROUP - - -def _inputs(bs, taps, dtype, seed=0): - g = torch.Generator().manual_seed(seed) - x = torch.randn(bs * BLOCK, HIDDEN, generator=g).to(dtype) - delta = (0.5 * torch.randn(bs * BLOCK, taps, NUM_GROUPS, generator=g)).to(dtype) - base = torch.randn(taps, HIDDEN, generator=g).to(dtype) - return x, delta, base - - -def _stance(compiled): - # The engine calls the torch.compile'd _grouped_conv (inductor here); the - # eager stance runs its original Python instead. - if compiled: - return contextlib.nullcontext() - return torch.compiler.set_stance("force_eager") - - -def _conv(x, delta, base, taps, compiled=False): - with _stance(compiled): - return _grouped_conv(x, delta, base, BLOCK, NUM_GROUPS, GROUP, taps) - - -def _multiplicative_mask_conv(x, delta, base, taps): - # The previous formulation, kept to show finite outputs are unchanged. - blocks = x.unflatten(-1, (NUM_GROUPS, GROUP)) - coefficients = base.view(1, taps, NUM_GROUPS, GROUP) + delta.unsqueeze(-1) - out = coefficients[:, 0] * blocks - position = torch.arange(x.shape[0]) % BLOCK - for tap in range(1, taps): - shifted = F.pad(blocks[:-tap], (0, 0, 0, 0, tap, 0)) - out = out + coefficients[:, tap] * shifted * (position >= tap).view(-1, 1, 1) - return out.flatten(-2) - - -class TestDFlashGroupedConv(CustomTestCase): - def _check_nonfinite_rows_stay_in_their_request(self, compiled): - for taps in (2, 3): - for bad in (float("nan"), float("inf"), float("-inf")): - for row in range(BLOCK): - with self.subTest(taps=taps, bad=bad, row=row): - x, delta, base = _inputs(bs=3, taps=taps, dtype=torch.bfloat16) - clean = _conv(x, delta, base, taps, compiled) - clean = clean.view(3, BLOCK, HIDDEN) - x[row, 5] = bad # request 0 - out = _conv(x, delta, base, taps, compiled) - out = out.view(3, BLOCK, HIDDEN) - torch.testing.assert_close(out[1:], clean[1:], rtol=0, atol=0) - bad_rows = (~torch.isfinite(out[0])).any(-1).nonzero() - expected = torch.arange(row, min(row + taps, BLOCK)) - self.assertEqual(bad_rows.flatten().tolist(), expected.tolist()) - - def _check_module_prepare_and_finish_do_not_leak(self, compiled): - torch.manual_seed(0) - conv = DFlashGroupedConv(HIDDEN, BLOCK, 2, GROUP).to(torch.bfloat16) - x = torch.randn(2 * BLOCK, HIDDEN).to(torch.bfloat16) - y = torch.randn(2 * BLOCK, HIDDEN).to(torch.bfloat16) - with torch.no_grad(), _stance(compiled): - ref_x, kernel = conv.prepare(x) - ref_y = conv.finish(y, kernel) - x_bad, y_bad = x.clone(), y.clone() - x_bad[BLOCK - 1] = float("nan") - y_bad[BLOCK - 1] = float("inf") - out_x, _ = conv.prepare(x_bad) - out_y = conv.finish(y_bad, kernel) - torch.testing.assert_close(out_x[BLOCK:], ref_x[BLOCK:], rtol=0, atol=0) - torch.testing.assert_close(out_y[BLOCK:], ref_y[BLOCK:], rtol=0, atol=0) - - def test_nonfinite_rows_stay_in_their_request(self): - self._check_nonfinite_rows_stay_in_their_request(compiled=False) - - def test_compiled_nonfinite_rows_stay_in_their_request(self): - self._check_nonfinite_rows_stay_in_their_request(compiled=True) - - def test_module_prepare_and_finish_do_not_leak(self): - self._check_module_prepare_and_finish_do_not_leak(compiled=False) - - def test_compiled_module_prepare_and_finish_do_not_leak(self): - self._check_module_prepare_and_finish_do_not_leak(compiled=True) - - def test_finite_output_unchanged(self): - for dtype in (torch.float32, torch.bfloat16): - for taps in (2, 3): - for seed in range(4): - with self.subTest(dtype=dtype, taps=taps, seed=seed): - x, delta, base = _inputs(5, taps, dtype, seed) - # torch.equal treats -0.0 == 0.0: the old mask could - # produce -0.0 where the new one produces +0.0. - self.assertTrue( - torch.equal( - _conv(x, delta, base, taps), - _multiplicative_mask_conv(x, delta, base, taps), - ) - ) - - -if __name__ == "__main__": - unittest.main()