[AMD] Fix EAGLE3 speculative decoding with aiter attention backend (#19362)
This commit is contained in:
@@ -968,7 +968,6 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
elif forward_mode.is_target_verify():
|
elif forward_mode.is_target_verify():
|
||||||
if self.use_mla:
|
|
||||||
qo_indptr = self.qo_indptr[: bs + 1]
|
qo_indptr = self.qo_indptr[: bs + 1]
|
||||||
qo_indptr[: bs + 1] = torch.arange(
|
qo_indptr[: bs + 1] = torch.arange(
|
||||||
0,
|
0,
|
||||||
@@ -977,13 +976,17 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
|
if self.use_mla:
|
||||||
|
kv_lens = seq_lens + self.num_draft_tokens
|
||||||
|
else:
|
||||||
|
kv_lens = seq_lens
|
||||||
kv_indptr = self.kv_indptr[: bs + 1]
|
kv_indptr = self.kv_indptr[: bs + 1]
|
||||||
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
|
kv_indptr[1 : bs + 1] = torch.cumsum(kv_lens, dim=0)
|
||||||
kv_indices = self.cuda_graph_kv_indices
|
kv_indices = self.cuda_graph_kv_indices
|
||||||
create_flashinfer_kv_indices_triton[(bs,)](
|
create_flashinfer_kv_indices_triton[(bs,)](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
seq_lens,
|
kv_lens,
|
||||||
kv_indptr,
|
kv_indptr,
|
||||||
None,
|
None,
|
||||||
kv_indices,
|
kv_indices,
|
||||||
@@ -992,7 +995,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
||||||
max_q_len = self.num_draft_tokens
|
max_q_len = self.num_draft_tokens
|
||||||
|
|
||||||
# if self.kv_cache_dtype == fp8_dtype:
|
if self.use_mla:
|
||||||
if _use_mla_ps_kernel:
|
if _use_mla_ps_kernel:
|
||||||
|
|
||||||
num_kv_splits = self.max_split_per_batch
|
num_kv_splits = self.max_split_per_batch
|
||||||
@@ -1035,37 +1038,11 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_final_map=reduce_final_map,
|
reduce_final_map=reduce_final_map,
|
||||||
reduce_partial_map=reduce_partial_map,
|
reduce_partial_map=reduce_partial_map,
|
||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
# num_kv_splits_indptr=num_kv_splits_indptr,
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# Non-MLA target_verify cuda graph: use triton extend kernel metadata
|
|
||||||
draft_num = self.num_draft_tokens
|
|
||||||
qo_indptr = self.qo_indptr[: bs + 1]
|
|
||||||
qo_indptr[: bs + 1] = torch.arange(
|
|
||||||
0,
|
|
||||||
(1 + bs) * draft_num,
|
|
||||||
step=draft_num,
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=self.device,
|
|
||||||
)
|
|
||||||
|
|
||||||
kv_indptr = self.kv_indptr[: bs + 1]
|
|
||||||
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
|
|
||||||
|
|
||||||
kv_indices = self.cuda_graph_kv_indices
|
|
||||||
create_flashinfer_kv_indices_triton[(bs,)](
|
|
||||||
self.req_to_token,
|
|
||||||
req_pool_indices,
|
|
||||||
seq_lens,
|
|
||||||
kv_indptr,
|
|
||||||
None,
|
|
||||||
kv_indices,
|
|
||||||
self.req_to_token.stride(0),
|
|
||||||
)
|
|
||||||
|
|
||||||
custom_mask = self.cuda_graph_custom_mask
|
custom_mask = self.cuda_graph_custom_mask
|
||||||
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask
|
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask
|
||||||
seq_mask_len = draft_num * (seq_lens + draft_num)
|
seq_mask_len = max_q_len * (seq_lens + max_q_len)
|
||||||
mask_indptr = self.mask_indptr
|
mask_indptr = self.mask_indptr
|
||||||
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0)
|
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0)
|
||||||
mask_indptr = mask_indptr[: bs + 1]
|
mask_indptr = mask_indptr[: bs + 1]
|
||||||
@@ -1074,12 +1051,12 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_indptr,
|
kv_indptr,
|
||||||
kv_indices,
|
kv_indices,
|
||||||
qo_indptr,
|
qo_indptr,
|
||||||
None,
|
kv_last_page_len,
|
||||||
draft_num,
|
max_q_len,
|
||||||
None,
|
kv_indptr[-1].item(),
|
||||||
custom_mask=custom_mask,
|
custom_mask=custom_mask,
|
||||||
mask_indptr=mask_indptr,
|
mask_indptr=mask_indptr,
|
||||||
max_extend_len=draft_num,
|
max_extend_len=max_q_len,
|
||||||
)
|
)
|
||||||
elif forward_mode.is_draft_extend():
|
elif forward_mode.is_draft_extend():
|
||||||
num_tokens_per_bs = self.speculative_num_steps + 1
|
num_tokens_per_bs = self.speculative_num_steps + 1
|
||||||
@@ -1290,20 +1267,10 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
)
|
)
|
||||||
if not self.use_mla:
|
|
||||||
# Non-MLA: update custom_mask and mask_indptr for triton extend kernel
|
|
||||||
custom_mask = self.cuda_graph_custom_mask
|
|
||||||
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask
|
|
||||||
seq_mask_len = self.num_draft_tokens * (
|
|
||||||
seq_lens + self.num_draft_tokens
|
|
||||||
)
|
|
||||||
mask_indptr = self.mask_indptr[: bs + 1]
|
|
||||||
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0)
|
|
||||||
|
|
||||||
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
||||||
max_q_len = self.num_draft_tokens
|
max_q_len = self.num_draft_tokens
|
||||||
|
|
||||||
# if self.kv_cache_dtype == fp8_dtype:
|
if self.use_mla:
|
||||||
if _use_mla_ps_kernel:
|
if _use_mla_ps_kernel:
|
||||||
|
|
||||||
num_kv_splits = self.max_split_per_batch
|
num_kv_splits = self.max_split_per_batch
|
||||||
@@ -1346,7 +1313,24 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_final_map=reduce_final_map,
|
reduce_final_map=reduce_final_map,
|
||||||
reduce_partial_map=reduce_partial_map,
|
reduce_partial_map=reduce_partial_map,
|
||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
# num_kv_splits_indptr=num_kv_splits_indptr,
|
)
|
||||||
|
else:
|
||||||
|
custom_mask = self.cuda_graph_custom_mask
|
||||||
|
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask
|
||||||
|
seq_mask_len = max_q_len * (seq_lens + max_q_len)
|
||||||
|
mask_indptr = self.mask_indptr[: bs + 1]
|
||||||
|
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0)
|
||||||
|
|
||||||
|
self.forward_metadata = ForwardMetadata(
|
||||||
|
kv_indptr,
|
||||||
|
kv_indices,
|
||||||
|
qo_indptr,
|
||||||
|
kv_last_page_len,
|
||||||
|
max_q_len,
|
||||||
|
kv_indptr[-1].item(),
|
||||||
|
custom_mask=custom_mask,
|
||||||
|
mask_indptr=mask_indptr,
|
||||||
|
max_extend_len=max_q_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
elif forward_mode.is_draft_extend():
|
elif forward_mode.is_draft_extend():
|
||||||
@@ -1371,7 +1355,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
||||||
max_q_len = num_tokens_per_bs
|
max_q_len = num_tokens_per_bs
|
||||||
|
|
||||||
if _use_mla_ps_kernel:
|
if self.use_mla and _use_mla_ps_kernel:
|
||||||
|
|
||||||
num_kv_splits = self.max_split_per_batch
|
num_kv_splits = self.max_split_per_batch
|
||||||
|
|
||||||
@@ -1413,7 +1397,6 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_final_map=reduce_final_map,
|
reduce_final_map=reduce_final_map,
|
||||||
reduce_partial_map=reduce_partial_map,
|
reduce_partial_map=reduce_partial_map,
|
||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
# num_kv_splits_indptr=num_kv_splits_indptr,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -3,7 +3,8 @@ from types import SimpleNamespace
|
|||||||
|
|
||||||
import requests
|
import requests
|
||||||
|
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.srt.utils import is_hip
|
||||||
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
from sglang.test.run_eval import run_eval
|
from sglang.test.run_eval import run_eval
|
||||||
from sglang.test.server_fixtures.eagle_fixture import EagleServerBase
|
from sglang.test.server_fixtures.eagle_fixture import EagleServerBase
|
||||||
from sglang.test.test_utils import (
|
from sglang.test.test_utils import (
|
||||||
@@ -12,6 +13,9 @@ from sglang.test.test_utils import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=50, suite="stage-b-test-small-1-gpu")
|
register_cuda_ci(est_time=50, suite="stage-b-test-small-1-gpu")
|
||||||
|
register_amd_ci(est_time=50, suite="stage-b-test-small-1-gpu")
|
||||||
|
|
||||||
|
_is_hip = is_hip()
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3Basic(EagleServerBase):
|
class TestEagle3Basic(EagleServerBase):
|
||||||
@@ -22,7 +26,17 @@ class TestEagle3Basic(EagleServerBase):
|
|||||||
spec_steps = 2
|
spec_steps = 2
|
||||||
spec_topk = 1
|
spec_topk = 1
|
||||||
spec_tokens = 3
|
spec_tokens = 3
|
||||||
extra_args = ["--dtype=float16", "--chunked-prefill-size", 1024]
|
extra_args = (
|
||||||
|
[
|
||||||
|
"--dtype=float16",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
1024,
|
||||||
|
"--attention-backend",
|
||||||
|
"aiter",
|
||||||
|
]
|
||||||
|
if _is_hip
|
||||||
|
else ["--dtype=float16", "--chunked-prefill-size", 1024]
|
||||||
|
)
|
||||||
|
|
||||||
def test_mmlu(self):
|
def test_mmlu(self):
|
||||||
"""Override to add EAGLE-specific assertions"""
|
"""Override to add EAGLE-specific assertions"""
|
||||||
@@ -42,6 +56,9 @@ class TestEagle3Basic(EagleServerBase):
|
|||||||
"avg_spec_accept_length"
|
"avg_spec_accept_length"
|
||||||
]
|
]
|
||||||
print(f"{avg_spec_accept_length=}")
|
print(f"{avg_spec_accept_length=}")
|
||||||
|
if _is_hip:
|
||||||
|
self.assertGreater(avg_spec_accept_length, 2.24)
|
||||||
|
else:
|
||||||
self.assertGreater(avg_spec_accept_length, 2.26)
|
self.assertGreater(avg_spec_accept_length, 2.26)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user