Skip to content
Merged
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
27 changes: 21 additions & 6 deletions deepmd/pt_expt/kernels/triton/sezm/sweep_tile_configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,6 +152,7 @@
_flash_bwd_block_kernel,
_flash_bwd_kernel,
_flash_bwd_op,
_rotation_strides,
)
from deepmd.pt_expt.kernels.triton.sezm.so2_stack_fp16x3 import (
_mixing_stack_fp16x3_bwd_op,
Expand Down Expand Up @@ -1056,6 +1057,11 @@ def sweep_flash_bwd(
grad_pre_gate = torch.randn(n_nodes, dim, c_wide, device=device)
x_local = torch.randn(n_edge, n_focus, reduced_dim, cf, device=device)
wigner_dt = _block_diag_wigner(n_edge, lmax, device)
# The backward kernels take the rotation layout (``PACKED``) and the
# Wigner strides explicitly, exactly as ``_launch_backward`` passes them;
# the sweep always builds a dense ``(E, D, D)`` Wigner, so ``packed`` is
# False here and the strides are the plain contiguous ones.
packed, dt_se, dt_sr, dt_sk = _rotation_strides(wigner_dt)
rescale = torch.rand(dim, device=device, dtype=torch.float32) + 0.5
alpha = torch.rand(n_edge, n_focus, n_head, device=device, dtype=torch.float32)
dst = torch.randint(0, n_nodes, (n_edge,), device=device)
Expand Down Expand Up @@ -1086,6 +1092,7 @@ def launch_edge(warps: int, stages: int) -> tuple[torch.Tensor, ...]:
gxl = torch.empty_like(x_local)
gdt = torch.zeros_like(wigner_dt)
gw = torch.empty_like(alpha)
_, gdt_se, gdt_sr, gdt_sk = _rotation_strides(gdt)
wrap_triton(_flash_bwd_kernel.fn)[(n_edge,)](
grad_pre_gate,
x_local,
Expand All @@ -1105,19 +1112,19 @@ def launch_edge(warps: int, stages: int) -> tuple[torch.Tensor, ...]:
x_local.stride(1),
x_local.stride(2),
x_local.stride(3),
wigner_dt.stride(0),
wigner_dt.stride(1),
wigner_dt.stride(2),
dt_se,
dt_sr,
dt_sk,
alpha.stride(0),
alpha.stride(1),
alpha.stride(2),
gxl.stride(0),
gxl.stride(1),
gxl.stride(2),
gxl.stride(3),
gdt.stride(0),
gdt.stride(1),
gdt.stride(2),
gdt_se,
gdt_sr,
gdt_sk,
gw.stride(0),
gw.stride(1),
gw.stride(2),
Expand All @@ -1127,6 +1134,7 @@ def launch_edge(warps: int, stages: int) -> tuple[torch.Tensor, ...]:
NFOCUS=n_focus,
NHEAD=n_head,
BLOCK_C=triton.next_power_of_2(c_wide),
PACKED=packed,
num_warps=warps,
num_stages=stages,
)
Expand Down Expand Up @@ -1163,6 +1171,7 @@ def launch(block_e: int, warps: int, stages: int) -> tuple[torch.Tensor, ...]:
gxl = torch.empty_like(x_local)
gdt = torch.zeros_like(wigner_dt)
gw = torch.empty_like(alpha)
_, gdt_se, gdt_sr, gdt_sk = _rotation_strides(gdt)
wrap_triton(_flash_bwd_block_kernel)[(triton.cdiv(n_edge, block_e),)](
grad_pre_gate,
x_local,
Expand All @@ -1180,10 +1189,16 @@ def launch(block_e: int, warps: int, stages: int) -> tuple[torch.Tensor, ...]:
x_local.stride(1),
x_local.stride(2),
x_local.stride(3),
dt_se,
dt_sr,
dt_sk,
gxl.stride(0),
gxl.stride(1),
gxl.stride(2),
gxl.stride(3),
gdt_se,
gdt_sr,
gdt_sk,
L=lmax,
CF=cf,
CW=c_wide,
Expand Down
43 changes: 43 additions & 0 deletions source/tests/pt/model/test_descriptor_sezm_triton.py
Original file line number Diff line number Diff line change
Expand Up @@ -1970,6 +1970,49 @@ def fake_flash_sweep(cf, lmax, **kwargs):
)
self.assertEqual(calls, [])

def test_tune_missing_configs_runs_the_real_flash_bwd_sweep(self):
"""The flash-backward sweep launches the kernels it tunes.

The other tuning tests replace every sweep with a fake, so a sweep
launch that drifts from its kernel signature (as happened when the
``PACKED`` layout flag and the Wigner strides were added to the
flash-attention backward kernels) is only caught by running the real
sweep once. The built-in tables are hidden so the key is uncovered on
every GPU, and the edge count is kept tiny because the launch contract,
not the timing, is under test; the winners are therefore only checked
for membership in the candidate sets.
"""
import itertools
from unittest import (
mock,
)

from deepmd.pt_expt.kernels.triton.sezm import (
sweep_tile_configs,
)

tc = self.tile_configs
specs = {"flash_bwd": sweep_tile_configs._SWEEP_SPECS["flash_bwd"]}
with (
mock.patch.object(tc, "_builtin_tables", return_value={}),
mock.patch.dict(sweep_tile_configs._SWEEP_SPECS, specs, clear=True),
):
registered = sweep_tile_configs.tune_missing_configs(
[(32, 2, 1, 1)], level=2, device="cuda", n_edge=4096
)
key = (32, 2)
self.assertEqual(sorted(registered), ["flash_bwd_block", "flash_bwd_edge"])
self.assertIn(
registered["flash_bwd_edge"][key],
set(itertools.product((1, 2, 4), (1, 2))),
)
block = registered["flash_bwd_block"][key]
self.assertTrue(
block is None or block in set(sweep_tile_configs._EDGE_BLOCK_CANDIDATES)
)
self.assertTrue(tc.has_tile_config("flash_bwd_edge", key))
self.assertTrue(tc.has_tile_config("flash_bwd_block", key))


if __name__ == "__main__":
unittest.main()
Loading