[AMD][Spec] Enable GDN ReplaySSM target-verify on ROCm (#38184)

Co-authored-by: chuyeh <298092489+chuyeh@users.noreply.github.com>
This commit is contained in:
chuyeh
2026-09-17 01:29:14 -07:00
committed by GitHub
co-authored by chuyeh
parent acfde25d34
commit 2d08cc5ede
@@ -43,6 +43,11 @@ import torch
import triton
import triton.language as tl
from sglang.kernels.ops.attention.fla.utils import is_tf32_supported
from sglang.srt.utils import is_gfx95_supported
_IS_GFX95 = is_gfx95_supported()
@triton.jit
def gdn_replayssm_spec_circular_kernel(
@@ -1269,7 +1274,8 @@ def gdn_replayssm_spec_decode(
num_stages_flush: int = 2,
nk_flush: int = 2,
launch_mode: str = "both",
dot_precision: str = "tf32",
# Triton only accepts TF32 input precision on Ampere-or-newer NVIDIA GPUs.
dot_precision: str = "tf32" if is_tf32_supported else "ieee",
):
"""GDN cached speculative-decode on a CIRCULAR ring cache (split-qkv varlen).
@@ -1286,6 +1292,18 @@ def gdn_replayssm_spec_decode(
"""
if scale is None:
scale = checkpoint_state.shape[-1] ** -0.5
batch_size = query_start_loc.shape[0] - 1
# At the Qwen3.5 MTP shape, one wide V tile avoids duplicating q/k and gate
# work across two programs on gfx950.
if (
_IS_GFX95
and batch_size >= 32
and max_cache_len == 16
and max_spec_len == 4
and q.shape[-1] == 128
and v.shape[-1] == 128
):
block_v, num_stages = 128, 1
if is_flush.dtype != torch.int8:
is_flush = is_flush.to(torch.int8)