diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 65e5af5ef..5d29b16d7 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -968,31 +968,34 @@ class AiterAttnBackend(AttentionBackend): ) elif forward_mode.is_target_verify(): + qo_indptr = self.qo_indptr[: bs + 1] + qo_indptr[: bs + 1] = torch.arange( + 0, + (1 + bs) * self.num_draft_tokens, + step=self.num_draft_tokens, + dtype=torch.int32, + device=self.device, + ) if self.use_mla: - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[: bs + 1] = torch.arange( - 0, - (1 + bs) * self.num_draft_tokens, - step=self.num_draft_tokens, - 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), - ) - kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs] - max_q_len = self.num_draft_tokens + kv_lens = seq_lens + self.num_draft_tokens + else: + kv_lens = seq_lens + kv_indptr = self.kv_indptr[: bs + 1] + kv_indptr[1 : bs + 1] = torch.cumsum(kv_lens, dim=0) + kv_indices = self.cuda_graph_kv_indices + create_flashinfer_kv_indices_triton[(bs,)]( + self.req_to_token, + req_pool_indices, + kv_lens, + kv_indptr, + None, + kv_indices, + self.req_to_token.stride(0), + ) + kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs] + max_q_len = self.num_draft_tokens - # if self.kv_cache_dtype == fp8_dtype: + if self.use_mla: if _use_mla_ps_kernel: num_kv_splits = self.max_split_per_batch @@ -1035,37 +1038,11 @@ class AiterAttnBackend(AttentionBackend): reduce_final_map=reduce_final_map, reduce_partial_map=reduce_partial_map, num_kv_splits=num_kv_splits, - # num_kv_splits_indptr=num_kv_splits_indptr, ) 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[: 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[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0) mask_indptr = mask_indptr[: bs + 1] @@ -1074,12 +1051,12 @@ class AiterAttnBackend(AttentionBackend): kv_indptr, kv_indices, qo_indptr, - None, - draft_num, - None, + kv_last_page_len, + max_q_len, + kv_indptr[-1].item(), custom_mask=custom_mask, mask_indptr=mask_indptr, - max_extend_len=draft_num, + max_extend_len=max_q_len, ) elif forward_mode.is_draft_extend(): num_tokens_per_bs = self.speculative_num_steps + 1 @@ -1290,64 +1267,71 @@ class AiterAttnBackend(AttentionBackend): kv_indices, 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] max_q_len = self.num_draft_tokens - # if self.kv_cache_dtype == fp8_dtype: - if _use_mla_ps_kernel: + if self.use_mla: + if _use_mla_ps_kernel: - num_kv_splits = self.max_split_per_batch + num_kv_splits = self.max_split_per_batch - self.make_mla_meta_data( - qo_indptr, + self.make_mla_meta_data( + qo_indptr, + kv_indptr, + kv_last_page_len, + self.work_metadata, + self.work_info_set, + self.work_indptr, + self.reduce_indptr, + self.reduce_final_map, + self.reduce_partial_map, + max_q_len, + fast_mode=fast_mode, + max_split_per_batch=num_kv_splits, + intra_batch_mode=intra_batch_mode, + ) + + work_metadata = self.work_metadata + work_info_set = self.work_info_set + work_indptr = self.work_indptr + + reduce_indptr = self.reduce_indptr + reduce_final_map = self.reduce_final_map + reduce_partial_map = self.reduce_partial_map + + self.forward_metadata = ForwardMetadata( kv_indptr, + kv_indices, + qo_indptr, kv_last_page_len, - self.work_metadata, - self.work_info_set, - self.work_indptr, - self.reduce_indptr, - self.reduce_final_map, - self.reduce_partial_map, max_q_len, - fast_mode=fast_mode, - max_split_per_batch=num_kv_splits, - intra_batch_mode=intra_batch_mode, + kv_indptr[-1].item(), + work_metadata=work_metadata, + work_info_set=work_info_set, + work_indptr=work_indptr, + reduce_indptr=reduce_indptr, + reduce_final_map=reduce_final_map, + reduce_partial_map=reduce_partial_map, + num_kv_splits=num_kv_splits, ) + 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) - work_metadata = self.work_metadata - work_info_set = self.work_info_set - work_indptr = self.work_indptr - - reduce_indptr = self.reduce_indptr - reduce_final_map = self.reduce_final_map - reduce_partial_map = self.reduce_partial_map - - self.forward_metadata = ForwardMetadata( - kv_indptr, - kv_indices, - qo_indptr, - kv_last_page_len, - max_q_len, - kv_indptr[-1].item(), - work_metadata=work_metadata, - work_info_set=work_info_set, - work_indptr=work_indptr, - reduce_indptr=reduce_indptr, - reduce_final_map=reduce_final_map, - reduce_partial_map=reduce_partial_map, - num_kv_splits=num_kv_splits, - # num_kv_splits_indptr=num_kv_splits_indptr, - ) + 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(): num_tokens_per_bs = self.speculative_num_steps + 1 @@ -1371,7 +1355,7 @@ class AiterAttnBackend(AttentionBackend): kv_last_page_len = self.cuda_graph_kv_last_page_len[: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 @@ -1413,7 +1397,6 @@ class AiterAttnBackend(AttentionBackend): reduce_final_map=reduce_final_map, reduce_partial_map=reduce_partial_map, num_kv_splits=num_kv_splits, - # num_kv_splits_indptr=num_kv_splits_indptr, ) else: diff --git a/test/registered/spec/eagle/test_eagle3_basic.py b/test/registered/spec/eagle/test_eagle3_basic.py index b60b167d0..98526c7e4 100644 --- a/test/registered/spec/eagle/test_eagle3_basic.py +++ b/test/registered/spec/eagle/test_eagle3_basic.py @@ -3,7 +3,8 @@ from types import SimpleNamespace 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.server_fixtures.eagle_fixture import EagleServerBase 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_amd_ci(est_time=50, suite="stage-b-test-small-1-gpu") + +_is_hip = is_hip() class TestEagle3Basic(EagleServerBase): @@ -22,7 +26,17 @@ class TestEagle3Basic(EagleServerBase): spec_steps = 2 spec_topk = 1 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): """Override to add EAGLE-specific assertions""" @@ -42,7 +56,10 @@ class TestEagle3Basic(EagleServerBase): "avg_spec_accept_length" ] print(f"{avg_spec_accept_length=}") - self.assertGreater(avg_spec_accept_length, 2.26) + if _is_hip: + self.assertGreater(avg_spec_accept_length, 2.24) + else: + self.assertGreater(avg_spec_accept_length, 2.26) if __name__ == "__main__":