gpt-oss decode performance optimization (#20392)
Co-authored-by: wunhuang <wunhuang@amd.com>
This commit is contained in:
@@ -13,7 +13,10 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
|
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
|
from sglang.srt.layers.attention.utils import (
|
||||||
|
create_flashinfer_kv_indices_triton,
|
||||||
|
create_flashmla_kv_indices_triton,
|
||||||
|
)
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
get_attention_tp_size,
|
get_attention_tp_size,
|
||||||
is_dp_attention_enabled,
|
is_dp_attention_enabled,
|
||||||
@@ -39,14 +42,19 @@ try:
|
|||||||
paged_attention_ragged,
|
paged_attention_ragged,
|
||||||
)
|
)
|
||||||
from aiter.mla import mla_decode_fwd, mla_prefill_fwd
|
from aiter.mla import mla_decode_fwd, mla_prefill_fwd
|
||||||
|
from aiter.ops.triton.attention.unified_attention import unified_attention
|
||||||
except ImportError:
|
except ImportError:
|
||||||
print(
|
print(
|
||||||
"aiter is AMD specific kernel library. Please make sure aiter is installed on your AMD device."
|
"aiter is AMD specific kernel library. Please make sure aiter is installed on your AMD device."
|
||||||
)
|
)
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import AttentionArch
|
from sglang.srt.configs.model_config import AttentionArch
|
||||||
from sglang.srt.layers.attention.utils import pad_sequence_with_mask
|
from sglang.srt.layers.attention.utils import (
|
||||||
|
launch_reshape_and_cache_flash,
|
||||||
|
pad_sequence_with_mask,
|
||||||
|
)
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
||||||
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.utils import get_bool_env_var
|
from sglang.srt.utils import get_bool_env_var
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -93,6 +101,7 @@ class ForwardMetadata:
|
|||||||
mask_indptr: Optional[torch.Tensor] = None
|
mask_indptr: Optional[torch.Tensor] = None
|
||||||
max_extend_len: Optional[int] = None
|
max_extend_len: Optional[int] = None
|
||||||
fp8_prefill_kv_indices: Optional[torch.Tensor] = None
|
fp8_prefill_kv_indices: Optional[torch.Tensor] = None
|
||||||
|
swa_page_table: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
|
||||||
global_workspace_buffer = None
|
global_workspace_buffer = None
|
||||||
@@ -185,6 +194,18 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
model_runner, self
|
model_runner, self
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# sliding window attention
|
||||||
|
self.use_sliding_window_kv_pool = (
|
||||||
|
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
|
||||||
|
and model_runner.token_to_kv_pool.swa_layer_nums > 0
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.use_sliding_window_kv_pool:
|
||||||
|
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||||
|
self.use_triton_unified_attention = True
|
||||||
|
else:
|
||||||
|
self.use_triton_unified_attention = False
|
||||||
|
|
||||||
# aiter kernel related initialization
|
# aiter kernel related initialization
|
||||||
self.max_num_partitions = (
|
self.max_num_partitions = (
|
||||||
self.max_context_len + _AITER_PARTITION_SIZE_ROCM - 1
|
self.max_context_len + _AITER_PARTITION_SIZE_ROCM - 1
|
||||||
@@ -192,7 +213,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
|
|
||||||
nbyes_per_qo_elem = torch.finfo(torch.float32).bits // 8
|
nbyes_per_qo_elem = torch.finfo(torch.float32).bits // 8
|
||||||
|
|
||||||
if not self.use_mla:
|
if not (self.use_mla or self.use_triton_unified_attention):
|
||||||
self.workspace_buffer = torch.empty(
|
self.workspace_buffer = torch.empty(
|
||||||
(max_bs * self.num_head * self.max_num_partitions * self.head_dim)
|
(max_bs * self.num_head * self.max_num_partitions * self.head_dim)
|
||||||
* nbyes_per_qo_elem
|
* nbyes_per_qo_elem
|
||||||
@@ -439,6 +460,17 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
is_causal=is_causal,
|
is_causal=is_causal,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# for page size > 1 useful conversion function
|
||||||
|
def _transform_table_1_to_real(self, page_table: torch.Tensor) -> torch.Tensor:
|
||||||
|
page_size = self.page_size
|
||||||
|
if page_size == 1:
|
||||||
|
return page_table
|
||||||
|
max_seqlen_k = page_table.shape[1]
|
||||||
|
strided_indices = torch.arange(
|
||||||
|
0, max_seqlen_k, page_size, device=page_table.device, dtype=torch.int32
|
||||||
|
)
|
||||||
|
return page_table[:, strided_indices] // page_size
|
||||||
|
|
||||||
def _resolve_v2_num_draft_tokens(
|
def _resolve_v2_num_draft_tokens(
|
||||||
self,
|
self,
|
||||||
extend_seq_lens: Optional[torch.Tensor] = None,
|
extend_seq_lens: Optional[torch.Tensor] = None,
|
||||||
@@ -591,6 +623,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
qo_indptr = None
|
qo_indptr = None
|
||||||
kv_last_page_len = None
|
kv_last_page_len = None
|
||||||
max_q_len = None
|
max_q_len = None
|
||||||
|
max_kv_len = None
|
||||||
|
|
||||||
work_metadata = None
|
work_metadata = None
|
||||||
work_indptr = None
|
work_indptr = None
|
||||||
@@ -600,12 +633,14 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_partial_map = None
|
reduce_partial_map = None
|
||||||
|
|
||||||
num_kv_splits = None
|
num_kv_splits = None
|
||||||
# num_kv_splits_indptr = None
|
swa_page_table = None
|
||||||
|
|
||||||
if forward_batch.forward_mode.is_decode_or_idle():
|
if forward_batch.forward_mode.is_decode_or_idle():
|
||||||
if spec_info is None or forward_batch.forward_mode.is_idle():
|
if spec_info is None or forward_batch.forward_mode.is_idle():
|
||||||
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
|
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
|
||||||
kv_indptr = kv_indptr[: bs + 1]
|
kv_indptr = kv_indptr[: bs + 1]
|
||||||
|
|
||||||
|
if not self.use_triton_unified_attention:
|
||||||
kv_indices = self._get_kv_indices_scratch(
|
kv_indices = self._get_kv_indices_scratch(
|
||||||
forward_batch.seq_lens_sum, forward_batch.seq_lens.device
|
forward_batch.seq_lens_sum, forward_batch.seq_lens.device
|
||||||
)
|
)
|
||||||
@@ -618,6 +653,41 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
max_q_len = 1
|
||||||
|
page_size = self.page_size
|
||||||
|
max_kv_len = torch.max(forward_batch.seq_lens).item()
|
||||||
|
max_num_blocks_per_seq = (max_kv_len + page_size - 1) // page_size
|
||||||
|
kv_indices = torch.zeros(
|
||||||
|
bs, max_kv_len, dtype=torch.int32, device=self.device
|
||||||
|
)
|
||||||
|
|
||||||
|
create_flashmla_kv_indices_triton[(bs,)](
|
||||||
|
self.req_to_token,
|
||||||
|
forward_batch.req_pool_indices,
|
||||||
|
forward_batch.seq_lens,
|
||||||
|
None,
|
||||||
|
kv_indices,
|
||||||
|
self.req_to_token.stride(0),
|
||||||
|
max_kv_len,
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.use_sliding_window_kv_pool:
|
||||||
|
swa_page_table = (
|
||||||
|
self.token_to_kv_pool.translate_loc_from_full_to_swa(
|
||||||
|
kv_indices
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
kv_indices = self._transform_table_1_to_real(kv_indices)
|
||||||
|
swa_page_table = self._transform_table_1_to_real(swa_page_table)
|
||||||
|
|
||||||
|
qo_indptr = self.qo_indptr[: bs + 1]
|
||||||
|
qo_indptr[1 : bs + 1] = torch.cumsum(
|
||||||
|
self.kv_last_page_len[:bs], dim=0
|
||||||
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
|
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
|
||||||
bs = kv_indptr.shape[0] - 1
|
bs = kv_indptr.shape[0] - 1
|
||||||
@@ -662,7 +732,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
qo_indptr,
|
qo_indptr,
|
||||||
kv_last_page_len,
|
kv_last_page_len,
|
||||||
max_q_len,
|
max_q_len,
|
||||||
None,
|
max_kv_len,
|
||||||
work_metadata=work_metadata,
|
work_metadata=work_metadata,
|
||||||
work_info_set=work_info_set,
|
work_info_set=work_info_set,
|
||||||
work_indptr=work_indptr,
|
work_indptr=work_indptr,
|
||||||
@@ -671,6 +741,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_partial_map=reduce_partial_map,
|
reduce_partial_map=reduce_partial_map,
|
||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
run_graph=False,
|
run_graph=False,
|
||||||
|
swa_page_table=swa_page_table,
|
||||||
)
|
)
|
||||||
|
|
||||||
elif forward_batch.forward_mode.is_draft_extend_v2():
|
elif forward_batch.forward_mode.is_draft_extend_v2():
|
||||||
@@ -1054,6 +1125,14 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
encoder_lens=forward_batch.encoder_lens,
|
encoder_lens=forward_batch.encoder_lens,
|
||||||
spec_info=None,
|
spec_info=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if self.use_sliding_window_kv_pool:
|
||||||
|
swa_page_table = (
|
||||||
|
self.token_to_kv_pool.translate_loc_from_full_to_swa(
|
||||||
|
self.indices_updater_prefill.kv_indices
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
self.forward_metadata = ForwardMetadata(
|
self.forward_metadata = ForwardMetadata(
|
||||||
self.indices_updater_prefill.kv_indptr,
|
self.indices_updater_prefill.kv_indptr,
|
||||||
self.indices_updater_prefill.kv_indices,
|
self.indices_updater_prefill.kv_indices,
|
||||||
@@ -1061,6 +1140,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
None,
|
None,
|
||||||
self.indices_updater_prefill.max_q_len,
|
self.indices_updater_prefill.max_q_len,
|
||||||
self.indices_updater_prefill.max_kv_len,
|
self.indices_updater_prefill.max_kv_len,
|
||||||
|
swa_page_table=swa_page_table,
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_cuda_graph_state(
|
def init_cuda_graph_state(
|
||||||
@@ -1071,8 +1151,11 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
):
|
):
|
||||||
self.cuda_graph_kv_last_page_len = torch.ones(max_bs, dtype=torch.int)
|
self.cuda_graph_kv_last_page_len = torch.ones(max_bs, dtype=torch.int)
|
||||||
if kv_indices_buf is None:
|
if kv_indices_buf is None:
|
||||||
|
max_num_blocks_per_seq = (
|
||||||
|
self.max_context_len + self.page_size - 1
|
||||||
|
) // self.page_size
|
||||||
self.cuda_graph_kv_indices = torch.zeros(
|
self.cuda_graph_kv_indices = torch.zeros(
|
||||||
(max_bs * self.max_context_len),
|
(max_bs * max_num_blocks_per_seq),
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
@@ -1111,6 +1194,16 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
self.reduce_final_map = None
|
self.reduce_final_map = None
|
||||||
self.reduce_partial_map = None
|
self.reduce_partial_map = None
|
||||||
|
|
||||||
|
if self.use_sliding_window_kv_pool:
|
||||||
|
max_num_blocks_per_seq = (
|
||||||
|
self.max_context_len + self.page_size - 1
|
||||||
|
) // self.page_size
|
||||||
|
self.cuda_graph_swa_page_table = torch.zeros(
|
||||||
|
(max_bs, max_num_blocks_per_seq),
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
def init_forward_metadata_capture_cuda_graph(
|
def init_forward_metadata_capture_cuda_graph(
|
||||||
self,
|
self,
|
||||||
bs: int,
|
bs: int,
|
||||||
@@ -1133,12 +1226,18 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_final_map = None
|
reduce_final_map = None
|
||||||
reduce_partial_map = None
|
reduce_partial_map = None
|
||||||
|
|
||||||
|
swa_page_table = None
|
||||||
|
|
||||||
|
max_kv_len = torch.max(seq_lens).item()
|
||||||
|
|
||||||
if forward_mode.is_decode_or_idle():
|
if forward_mode.is_decode_or_idle():
|
||||||
qo_indptr = None
|
qo_indptr = None
|
||||||
kv_last_page_len = None
|
kv_last_page_len = None
|
||||||
max_q_len = None
|
max_q_len = None
|
||||||
|
|
||||||
if spec_info is None:
|
if spec_info is None:
|
||||||
|
|
||||||
|
if not self.use_triton_unified_attention:
|
||||||
kv_indptr = self.kv_indptr
|
kv_indptr = self.kv_indptr
|
||||||
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
|
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
|
||||||
kv_indptr = kv_indptr[: bs + 1]
|
kv_indptr = kv_indptr[: bs + 1]
|
||||||
@@ -1152,6 +1251,42 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
max_q_len = 1
|
||||||
|
max_num_blocks_per_seq = (
|
||||||
|
self.max_context_len + self.page_size - 1
|
||||||
|
) // self.page_size
|
||||||
|
kv_indices = self.cuda_graph_kv_indices.view(
|
||||||
|
-1, max_num_blocks_per_seq
|
||||||
|
)
|
||||||
|
|
||||||
|
page_indices = self.req_to_token[req_pool_indices[:bs], :max_kv_len]
|
||||||
|
|
||||||
|
if self.use_sliding_window_kv_pool:
|
||||||
|
swa_page_indices = (
|
||||||
|
self.token_to_kv_pool.translate_loc_from_full_to_swa(
|
||||||
|
page_indices
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
page_indices = self._transform_table_1_to_real(page_indices)
|
||||||
|
swa_page_indices = self._transform_table_1_to_real(
|
||||||
|
swa_page_indices
|
||||||
|
)
|
||||||
|
|
||||||
|
new_rows = swa_page_indices.shape[0]
|
||||||
|
new_cols = swa_page_indices.shape[1]
|
||||||
|
|
||||||
|
kv_indices[:new_rows, :new_cols].copy_(page_indices)
|
||||||
|
swa_page_table = self.cuda_graph_swa_page_table
|
||||||
|
swa_page_table[:new_rows, :new_cols].copy_(swa_page_indices)
|
||||||
|
|
||||||
|
qo_indptr = self.qo_indptr[: bs + 1]
|
||||||
|
qo_indptr[1 : bs + 1] = torch.cumsum(
|
||||||
|
self.cuda_graph_kv_last_page_len[:bs], dim=0
|
||||||
|
)
|
||||||
|
|
||||||
|
kv_indptr = None
|
||||||
else:
|
else:
|
||||||
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
|
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
|
||||||
|
|
||||||
@@ -1196,7 +1331,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
qo_indptr,
|
qo_indptr,
|
||||||
kv_last_page_len,
|
kv_last_page_len,
|
||||||
max_q_len,
|
max_q_len,
|
||||||
kv_indptr[-1].item(),
|
max_kv_len,
|
||||||
work_metadata=work_metadata,
|
work_metadata=work_metadata,
|
||||||
work_info_set=work_info_set,
|
work_info_set=work_info_set,
|
||||||
work_indptr=work_indptr,
|
work_indptr=work_indptr,
|
||||||
@@ -1204,6 +1339,7 @@ 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,
|
||||||
|
swa_page_table=swa_page_table,
|
||||||
)
|
)
|
||||||
|
|
||||||
elif forward_mode.is_target_verify():
|
elif forward_mode.is_target_verify():
|
||||||
@@ -1469,12 +1605,16 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_final_map = None
|
reduce_final_map = None
|
||||||
reduce_partial_map = None
|
reduce_partial_map = None
|
||||||
|
|
||||||
|
swa_page_table = None
|
||||||
|
max_kv_len = torch.max(seq_lens).item()
|
||||||
|
|
||||||
if forward_mode.is_decode_or_idle():
|
if forward_mode.is_decode_or_idle():
|
||||||
qo_indptr = None
|
qo_indptr = None
|
||||||
kv_last_page_len = None
|
kv_last_page_len = None
|
||||||
max_q_len = None
|
max_q_len = None
|
||||||
|
|
||||||
if spec_info is None:
|
if spec_info is None:
|
||||||
|
if not self.use_triton_unified_attention:
|
||||||
kv_indptr = self.kv_indptr
|
kv_indptr = self.kv_indptr
|
||||||
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
|
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
|
||||||
kv_indptr = kv_indptr[: bs + 1]
|
kv_indptr = kv_indptr[: bs + 1]
|
||||||
@@ -1488,6 +1628,42 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_indices,
|
kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
max_q_len = 1
|
||||||
|
max_num_blocks_per_seq = (
|
||||||
|
self.max_context_len + self.page_size - 1
|
||||||
|
) // self.page_size
|
||||||
|
kv_indices = self.cuda_graph_kv_indices.view(
|
||||||
|
-1, max_num_blocks_per_seq
|
||||||
|
)
|
||||||
|
|
||||||
|
page_indices = self.req_to_token[req_pool_indices[:bs], :max_kv_len]
|
||||||
|
|
||||||
|
if self.use_sliding_window_kv_pool:
|
||||||
|
swa_page_indices = (
|
||||||
|
self.token_to_kv_pool.translate_loc_from_full_to_swa(
|
||||||
|
page_indices
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
page_indices = self._transform_table_1_to_real(page_indices)
|
||||||
|
swa_page_indices = self._transform_table_1_to_real(
|
||||||
|
swa_page_indices
|
||||||
|
)
|
||||||
|
|
||||||
|
new_rows = swa_page_indices.shape[0]
|
||||||
|
new_cols = swa_page_indices.shape[1]
|
||||||
|
|
||||||
|
kv_indices[:new_rows, :new_cols].copy_(page_indices)
|
||||||
|
swa_page_table = self.cuda_graph_swa_page_table
|
||||||
|
swa_page_table[:new_rows, :new_cols].copy_(swa_page_indices)
|
||||||
|
|
||||||
|
qo_indptr = self.qo_indptr[: bs + 1]
|
||||||
|
qo_indptr[1 : bs + 1] = torch.cumsum(
|
||||||
|
self.cuda_graph_kv_last_page_len[:bs], dim=0
|
||||||
|
)
|
||||||
|
|
||||||
|
kv_indptr = None
|
||||||
else:
|
else:
|
||||||
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
|
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
|
||||||
|
|
||||||
@@ -1532,7 +1708,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
qo_indptr,
|
qo_indptr,
|
||||||
kv_last_page_len,
|
kv_last_page_len,
|
||||||
max_q_len,
|
max_q_len,
|
||||||
kv_indptr[-1].item(),
|
max_kv_len,
|
||||||
work_metadata=work_metadata,
|
work_metadata=work_metadata,
|
||||||
work_info_set=work_info_set,
|
work_info_set=work_info_set,
|
||||||
work_indptr=work_indptr,
|
work_indptr=work_indptr,
|
||||||
@@ -1540,6 +1716,8 @@ 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,
|
||||||
|
swa_page_table=swa_page_table,
|
||||||
|
# num_kv_splits_indptr=num_kv_splits_indptr,
|
||||||
)
|
)
|
||||||
|
|
||||||
elif forward_mode.is_target_verify():
|
elif forward_mode.is_target_verify():
|
||||||
@@ -1794,18 +1972,42 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
save_kv_cache=True,
|
save_kv_cache=True,
|
||||||
sinks=None,
|
sinks=None,
|
||||||
):
|
):
|
||||||
|
self.logits_soft_cap = layer.logit_cap
|
||||||
|
|
||||||
cache_loc = (
|
cache_loc = (
|
||||||
forward_batch.out_cache_loc
|
forward_batch.out_cache_loc
|
||||||
if not layer.is_cross_attention
|
if not layer.is_cross_attention
|
||||||
else forward_batch.encoder_out_cache_loc
|
else forward_batch.encoder_out_cache_loc
|
||||||
)
|
)
|
||||||
|
|
||||||
self.logits_soft_cap = layer.logit_cap
|
|
||||||
|
|
||||||
if k is not None:
|
if k is not None:
|
||||||
assert v is not None
|
assert v is not None
|
||||||
if save_kv_cache:
|
if save_kv_cache:
|
||||||
if self.use_mla:
|
if self.use_triton_unified_attention:
|
||||||
|
token_to_kv_pool = forward_batch.token_to_kv_pool
|
||||||
|
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||||
|
layer.layer_id
|
||||||
|
)
|
||||||
|
slot_mapping_swa = token_to_kv_pool.full_to_swa_index_mapping
|
||||||
|
|
||||||
|
launch_reshape_and_cache_flash(
|
||||||
|
k.view(-1, layer.tp_k_head_num, layer.qk_head_dim),
|
||||||
|
v.view(-1, layer.tp_v_head_num, layer.v_head_dim),
|
||||||
|
k_cache.view(
|
||||||
|
-1, self.page_size, layer.tp_k_head_num, layer.qk_head_dim
|
||||||
|
),
|
||||||
|
v_cache.view(
|
||||||
|
-1, self.page_size, layer.tp_v_head_num, layer.v_head_dim
|
||||||
|
),
|
||||||
|
cache_loc,
|
||||||
|
(
|
||||||
|
slot_mapping_swa.long()
|
||||||
|
if layer.sliding_window_size > 0
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
elif self.use_mla:
|
||||||
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
|
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
|
||||||
else:
|
else:
|
||||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||||
@@ -2132,8 +2334,14 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
v_cache = v_cache.to(dtype)
|
v_cache = v_cache.to(dtype)
|
||||||
|
|
||||||
window_size = (-1, -1)
|
window_size = (-1, -1)
|
||||||
|
page_table = self.forward_metadata.kv_indices
|
||||||
|
|
||||||
if layer.sliding_window_size is not None and layer.sliding_window_size > -1:
|
if layer.sliding_window_size is not None and layer.sliding_window_size > -1:
|
||||||
window_size = (layer.sliding_window_size, -1)
|
window_size = (layer.sliding_window_size, -1)
|
||||||
|
# page_table = self.token_to_kv_pool.translate_loc_from_full_to_swa(
|
||||||
|
# page_table
|
||||||
|
# )
|
||||||
|
page_table = self.forward_metadata.swa_page_table
|
||||||
|
|
||||||
o = mha_batch_prefill_func(
|
o = mha_batch_prefill_func(
|
||||||
q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
|
q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
|
||||||
@@ -2141,7 +2349,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
v_cache,
|
v_cache,
|
||||||
self.qo_indptr[:bs0],
|
self.qo_indptr[:bs0],
|
||||||
self.forward_metadata.kv_indptr[:bs0],
|
self.forward_metadata.kv_indptr[:bs0],
|
||||||
self.forward_metadata.kv_indices,
|
page_table,
|
||||||
self.forward_metadata.max_q_len,
|
self.forward_metadata.max_q_len,
|
||||||
self.forward_metadata.max_kv_len,
|
self.forward_metadata.max_kv_len,
|
||||||
causal=True,
|
causal=True,
|
||||||
@@ -2163,6 +2371,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
layer: RadixAttention,
|
layer: RadixAttention,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
save_kv_cache=True,
|
save_kv_cache=True,
|
||||||
|
sinks=None,
|
||||||
):
|
):
|
||||||
|
|
||||||
q = q.reshape(-1, layer.tp_q_head_num * layer.qk_head_dim)
|
q = q.reshape(-1, layer.tp_q_head_num * layer.qk_head_dim)
|
||||||
@@ -2176,7 +2385,26 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
o = torch.empty_like(q, dtype=self.input_dtype)
|
o = torch.empty_like(q, dtype=self.input_dtype)
|
||||||
|
|
||||||
if save_kv_cache:
|
if save_kv_cache:
|
||||||
|
if self.use_triton_unified_attention:
|
||||||
|
token_to_kv_pool = forward_batch.token_to_kv_pool
|
||||||
|
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||||
|
layer.layer_id
|
||||||
|
)
|
||||||
|
slot_mapping_swa = token_to_kv_pool.full_to_swa_index_mapping
|
||||||
|
|
||||||
|
launch_reshape_and_cache_flash(
|
||||||
|
k.view(-1, layer.tp_k_head_num, layer.qk_head_dim),
|
||||||
|
v.view(-1, layer.tp_v_head_num, layer.v_head_dim),
|
||||||
|
k_cache.view(
|
||||||
|
-1, self.page_size, layer.tp_k_head_num, layer.qk_head_dim
|
||||||
|
),
|
||||||
|
v_cache.view(
|
||||||
|
-1, self.page_size, layer.tp_v_head_num, layer.v_head_dim
|
||||||
|
),
|
||||||
|
forward_batch.out_cache_loc,
|
||||||
|
slot_mapping_swa.long() if layer.sliding_window_size > 0 else None,
|
||||||
|
)
|
||||||
|
else:
|
||||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||||
layer, forward_batch.out_cache_loc, k, v
|
layer, forward_batch.out_cache_loc, k, v
|
||||||
)
|
)
|
||||||
@@ -2230,6 +2458,47 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
k_cache = k_cache.to(dtype)
|
k_cache = k_cache.to(dtype)
|
||||||
v_cache = v_cache.to(dtype)
|
v_cache = v_cache.to(dtype)
|
||||||
|
|
||||||
|
if self.use_triton_unified_attention:
|
||||||
|
|
||||||
|
bs = forward_batch.batch_size
|
||||||
|
window_size = (-1, -1)
|
||||||
|
page_table = self.forward_metadata.kv_indices
|
||||||
|
|
||||||
|
if (
|
||||||
|
layer.sliding_window_size is not None
|
||||||
|
and layer.sliding_window_size > -1
|
||||||
|
):
|
||||||
|
window_size = (layer.sliding_window_size - 1, 0)
|
||||||
|
page_table = self.forward_metadata.swa_page_table
|
||||||
|
|
||||||
|
o = torch.empty_like(q)
|
||||||
|
|
||||||
|
max_kv_len = page_table.shape[1]
|
||||||
|
|
||||||
|
unified_attention(
|
||||||
|
q=q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
|
||||||
|
k=k_cache.view(
|
||||||
|
-1, self.page_size, layer.tp_k_head_num, layer.qk_head_dim
|
||||||
|
),
|
||||||
|
v=v_cache.view(
|
||||||
|
-1, self.page_size, layer.tp_v_head_num, layer.v_head_dim
|
||||||
|
),
|
||||||
|
out=o.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
|
||||||
|
cu_seqlens_q=self.forward_metadata.qo_indptr,
|
||||||
|
seqused_k=forward_batch.seq_lens,
|
||||||
|
max_seqlen_q=self.forward_metadata.max_q_len,
|
||||||
|
max_seqlen_k=max_kv_len,
|
||||||
|
softmax_scale=self.scale,
|
||||||
|
causal=True,
|
||||||
|
window_size=window_size,
|
||||||
|
block_table=page_table,
|
||||||
|
softcap=0,
|
||||||
|
q_descale=None,
|
||||||
|
k_descale=None,
|
||||||
|
v_descale=None,
|
||||||
|
sinks=sinks,
|
||||||
|
)
|
||||||
|
else:
|
||||||
paged_attention_ragged(
|
paged_attention_ragged(
|
||||||
o.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
|
o.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
|
||||||
self.workspace_buffer,
|
self.workspace_buffer,
|
||||||
|
|||||||
@@ -472,3 +472,189 @@ def concat_mla_absorb_q_general(q_nope, q_rope):
|
|||||||
return concat_mla_absorb_q(q_nope, q_rope)
|
return concat_mla_absorb_q(q_nope, q_rope)
|
||||||
else:
|
else:
|
||||||
return torch.cat([q_nope, q_rope], dim=-1)
|
return torch.cat([q_nope, q_rope], dim=-1)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def reshape_and_cache_flash(
|
||||||
|
key_ptr,
|
||||||
|
value_ptr,
|
||||||
|
key_cache_ptr,
|
||||||
|
value_cache_ptr,
|
||||||
|
slot_mapping_ptr,
|
||||||
|
swa_slot_mapping_ptr,
|
||||||
|
k_scale_ptr,
|
||||||
|
v_scale_ptr,
|
||||||
|
block_stride,
|
||||||
|
key_stride,
|
||||||
|
value_stride,
|
||||||
|
num_heads,
|
||||||
|
head_size,
|
||||||
|
block_size,
|
||||||
|
HEAD_BLOCK: tl.constexpr,
|
||||||
|
BLOCK_D: tl.constexpr,
|
||||||
|
HAS_SWA: tl.constexpr,
|
||||||
|
USE_SCALE: tl.constexpr,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Triton kernel for reshaping per-token K/V tensors into paged KV cache layout.
|
||||||
|
|
||||||
|
Source layout:
|
||||||
|
key/value: [num_tokens, num_heads, head_size]
|
||||||
|
|
||||||
|
Target cache layout:
|
||||||
|
cache: [num_blocks, block_size, num_heads, head_size]
|
||||||
|
|
||||||
|
Each Triton program instance handles:
|
||||||
|
- one token (program_id(0))
|
||||||
|
- one block of heads (program_id(1))
|
||||||
|
|
||||||
|
Features:
|
||||||
|
- optional SWA slot remapping
|
||||||
|
- optional FP8 scale dequantization before cache write
|
||||||
|
|
||||||
|
Args:
|
||||||
|
key_ptr: Pointer to source key tensor.
|
||||||
|
value_ptr: Pointer to source value tensor.
|
||||||
|
key_cache_ptr: Pointer to destination key cache tensor.
|
||||||
|
value_cache_ptr: Pointer to destination value cache tensor.
|
||||||
|
slot_mapping_ptr: Maps token -> cache slot.
|
||||||
|
swa_slot_mapping_ptr: Optional second-stage slot remap for SWA mode.
|
||||||
|
k_scale_ptr: Optional key scaling factor pointer.
|
||||||
|
v_scale_ptr: Optional value scaling factor pointer.
|
||||||
|
block_stride: Stride between cache blocks.
|
||||||
|
key_stride: Stride between source key tokens.
|
||||||
|
value_stride: Stride between source value tokens.
|
||||||
|
num_heads: Number of attention heads.
|
||||||
|
head_size: Hidden dimension per head.
|
||||||
|
block_size: Number of slots per cache block.
|
||||||
|
HEAD_BLOCK: Number of heads processed per program.
|
||||||
|
BLOCK_D: Vectorized dimension size (power-of-2 padded).
|
||||||
|
HAS_SWA: Enable SWA remapping.
|
||||||
|
USE_SCALE: Enable scale division before storing.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# program ids
|
||||||
|
# pid0 = token
|
||||||
|
# pid1 = head block
|
||||||
|
# ----------------------------------
|
||||||
|
token_idx = tl.program_id(0)
|
||||||
|
head_block_idx = tl.program_id(1)
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# slot mapping
|
||||||
|
# ----------------------------------
|
||||||
|
slot_idx = tl.load(slot_mapping_ptr + token_idx)
|
||||||
|
|
||||||
|
if HAS_SWA:
|
||||||
|
slot_idx = tl.load(swa_slot_mapping_ptr + slot_idx)
|
||||||
|
|
||||||
|
if slot_idx < 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
block_idx = slot_idx // block_size
|
||||||
|
block_offset = slot_idx % block_size
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# head range
|
||||||
|
# ----------------------------------
|
||||||
|
head_idx = head_block_idx * HEAD_BLOCK + tl.arange(0, HEAD_BLOCK)
|
||||||
|
|
||||||
|
head_mask = head_idx < num_heads
|
||||||
|
|
||||||
|
dim_idx = tl.arange(0, BLOCK_D)
|
||||||
|
|
||||||
|
# shape = [HEAD_BLOCK, BLOCK_D]
|
||||||
|
offs = head_idx[:, None] * head_size + dim_idx[None, :]
|
||||||
|
|
||||||
|
mask = head_mask[:, None] & (dim_idx[None, :] < head_size)
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# source load
|
||||||
|
# ----------------------------------
|
||||||
|
src_key = token_idx * key_stride + offs
|
||||||
|
src_value = token_idx * value_stride + offs
|
||||||
|
|
||||||
|
k = tl.load(key_ptr + src_key, mask=mask)
|
||||||
|
v = tl.load(value_ptr + src_value, mask=mask)
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# optional scale
|
||||||
|
# ----------------------------------
|
||||||
|
if USE_SCALE:
|
||||||
|
k_scale = tl.load(k_scale_ptr)
|
||||||
|
v_scale = tl.load(v_scale_ptr)
|
||||||
|
|
||||||
|
k = k / k_scale
|
||||||
|
v = v / v_scale
|
||||||
|
|
||||||
|
# ----------------------------------
|
||||||
|
# target layout
|
||||||
|
# [block_idx, block_offset, head, dim]
|
||||||
|
# ----------------------------------
|
||||||
|
tgt = block_idx * block_stride + block_offset * num_heads * head_size + offs
|
||||||
|
|
||||||
|
tl.store(key_cache_ptr + tgt, k, mask=mask)
|
||||||
|
tl.store(value_cache_ptr + tgt, v, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def launch_reshape_and_cache_flash(
|
||||||
|
key,
|
||||||
|
value,
|
||||||
|
key_cache,
|
||||||
|
value_cache,
|
||||||
|
slot_mapping,
|
||||||
|
swa_slot_mapping=None,
|
||||||
|
k_scale=None,
|
||||||
|
v_scale=None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Launch wrapper for reshape_and_cache_flash Triton kernel.
|
||||||
|
|
||||||
|
This wrapper prepares launch configuration and dispatches the Triton kernel
|
||||||
|
that writes token-major K/V tensors into paged KV cache layout.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
key: Source key tensor [num_tokens, num_heads, head_size]
|
||||||
|
value: Source value tensor [num_tokens, num_heads, head_size]
|
||||||
|
key_cache: Destination key cache [num_blocks, block_size, num_heads, head_size]
|
||||||
|
value_cache: Destination value cache [num_blocks, block_size, num_heads, head_size]
|
||||||
|
slot_mapping: Token-to-cache slot mapping
|
||||||
|
swa_slot_mapping: Optional SWA remapping table
|
||||||
|
k_scale: Optional key scaling factor
|
||||||
|
v_scale: Optional value scaling factor
|
||||||
|
"""
|
||||||
|
|
||||||
|
num_tokens = key.shape[0]
|
||||||
|
num_heads = key.shape[1]
|
||||||
|
head_size = key.shape[2]
|
||||||
|
|
||||||
|
HEAD_BLOCK = 4
|
||||||
|
|
||||||
|
BLOCK_D = triton.next_power_of_2(head_size)
|
||||||
|
|
||||||
|
grid = (
|
||||||
|
num_tokens,
|
||||||
|
triton.cdiv(num_heads, HEAD_BLOCK),
|
||||||
|
)
|
||||||
|
|
||||||
|
reshape_and_cache_flash[grid](
|
||||||
|
key,
|
||||||
|
value,
|
||||||
|
key_cache,
|
||||||
|
value_cache,
|
||||||
|
slot_mapping,
|
||||||
|
swa_slot_mapping if swa_slot_mapping is not None else key,
|
||||||
|
k_scale if k_scale is not None else key,
|
||||||
|
v_scale if v_scale is not None else key,
|
||||||
|
key_cache.stride(0),
|
||||||
|
key.stride(0),
|
||||||
|
value.stride(0),
|
||||||
|
num_heads,
|
||||||
|
head_size,
|
||||||
|
key_cache.shape[1],
|
||||||
|
HEAD_BLOCK=HEAD_BLOCK,
|
||||||
|
BLOCK_D=BLOCK_D,
|
||||||
|
HAS_SWA=(swa_slot_mapping is not None),
|
||||||
|
USE_SCALE=(k_scale is not None),
|
||||||
|
)
|
||||||
|
|||||||
@@ -52,6 +52,7 @@ if _use_aiter:
|
|||||||
from aiter import ActivationType
|
from aiter import ActivationType
|
||||||
from aiter.fused_moe import fused_moe
|
from aiter.fused_moe import fused_moe
|
||||||
from aiter.ops.shuffle import shuffle_weight
|
from aiter.ops.shuffle import shuffle_weight
|
||||||
|
from aiter.tuned_gemm import tgemm
|
||||||
|
|
||||||
if _is_npu:
|
if _is_npu:
|
||||||
from sglang.srt.hardware_backend.npu.utils import npu_format_cast
|
from sglang.srt.hardware_backend.npu.utils import npu_format_cast
|
||||||
@@ -150,6 +151,9 @@ class UnquantizedLinearMethod(LinearMethodBase):
|
|||||||
output = output.view(x_shapes[0], x_shapes[1], -1)
|
output = output.view(x_shapes[0], x_shapes[1], -1)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
elif _use_aiter and type(layer.weight.data) is torch.Tensor:
|
||||||
|
return tgemm.mm(x, layer.weight, bias, otype=x.dtype)
|
||||||
|
|
||||||
return F.linear(x, layer.weight, bias)
|
return F.linear(x, layer.weight, bias)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1626,6 +1626,8 @@ class ServerArgs:
|
|||||||
self.attention_backend = "trtllm_mha"
|
self.attention_backend = "trtllm_mha"
|
||||||
elif is_sm90_supported():
|
elif is_sm90_supported():
|
||||||
self.attention_backend = "fa3"
|
self.attention_backend = "fa3"
|
||||||
|
elif is_hip():
|
||||||
|
self.attention_backend = "aiter"
|
||||||
else:
|
else:
|
||||||
self.attention_backend = "triton"
|
self.attention_backend = "triton"
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user