1672 lines
68 KiB
Python
1672 lines
68 KiB
Python
import math
|
|
from dataclasses import dataclass
|
|
from typing import ClassVar, TypeVar
|
|
|
|
import torch
|
|
import torch.distributed as dist
|
|
import torch.nn.functional as F
|
|
import torch_npu
|
|
from vllm.config import VllmConfig, get_current_vllm_config
|
|
from vllm.distributed import get_tp_group
|
|
from vllm.v1.attention.backend import AttentionCGSupport, AttentionMetadataBuilder
|
|
from vllm.v1.kv_cache_interface import AttentionSpec
|
|
|
|
from vllm_ascend.attention.abstract import DSAAttentionImpl
|
|
from vllm_ascend.attention.attention_v1 import AscendAttentionState
|
|
from vllm_ascend.attention.utils import (
|
|
AscendCommonAttentionMetadata,
|
|
maybe_save_kv_layer_to_connector,
|
|
notify_kv_cache_written,
|
|
split_decodes_and_prefills,
|
|
wait_for_kv_layer_from_connector,
|
|
)
|
|
from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec
|
|
from vllm_ascend.device.device_op import DeviceOperator
|
|
from vllm_ascend.distributed.utils import all_gather_async
|
|
from vllm_ascend.memcache_comm_fence import record_attention_compute_start
|
|
from vllm_ascend.ops.linear import AscendUnquantizedLinearMethod
|
|
from vllm_ascend.ops.rope_dsv4 import get_cos_and_sin_dsa, get_full_cos_and_sin_dsa
|
|
from vllm_ascend.quantization.methods.w8a8_dynamic import AscendW8A8DynamicLinearMethod
|
|
from vllm_ascend.utils import (
|
|
AscendDeviceType,
|
|
enable_dsa_cp_with_o_proj_tp,
|
|
get_ascend_device_type,
|
|
olora_tp_enable,
|
|
)
|
|
|
|
|
|
def hadamard_transform_ref(
|
|
x: torch.Tensor,
|
|
hadamard: torch.Tensor,
|
|
scale: float = 1.0, # type: ignore[assignment]
|
|
):
|
|
x_shape = x.shape
|
|
dim = x.shape[-1]
|
|
x = x.reshape(-1, dim)
|
|
log_dim = math.ceil(math.log2(dim))
|
|
dim_padded = 2**log_dim
|
|
if dim != dim_padded:
|
|
x = F.pad(x, (0, dim_padded - dim))
|
|
out = F.linear(x, hadamard)
|
|
out = out * scale
|
|
return out[..., :dim].reshape(*x_shape)
|
|
|
|
|
|
def rotate_activation(x: torch.Tensor, hadamard: torch.Tensor) -> torch.Tensor:
|
|
hidden_size = x.size(-1)
|
|
return hadamard_transform_ref(x, hadamard=hadamard, scale=hidden_size**-0.5)
|
|
|
|
|
|
def _has_prefill(attn_state: AscendAttentionState) -> bool:
|
|
return attn_state not in {
|
|
AscendAttentionState.DecodeOnly,
|
|
AscendAttentionState.SpecDecoding,
|
|
}
|
|
|
|
|
|
@dataclass
|
|
class DSACPMetadata:
|
|
"""Context-parallel metadata for sequence-sharded DSA execution."""
|
|
|
|
local_query_start_loc: torch.Tensor
|
|
local_seq_lens: torch.Tensor
|
|
local_start: int
|
|
local_end: int
|
|
tokens_per_rank: int
|
|
num_tokens_pad: int
|
|
local_sin: torch.Tensor = None
|
|
local_cos: torch.Tensor = None
|
|
|
|
|
|
@dataclass
|
|
class AscendDSAReqMetadata:
|
|
"""Unified per-request metadata — combines fields formerly split into
|
|
prefill and decode sub-structures.
|
|
|
|
All methods (builder, forward) operate on this single metadata,
|
|
without distinguishing prefill vs decode request types.
|
|
"""
|
|
|
|
input_positions: torch.Tensor
|
|
block_table: torch.Tensor
|
|
seq_lens: torch.Tensor
|
|
slot_mapping: torch.Tensor | None
|
|
block_size: int
|
|
query_start_loc: torch.Tensor
|
|
cp_metadata: DSACPMetadata
|
|
num_compressed_tokens: int | None = None
|
|
sin: torch.Tensor = None
|
|
cos: torch.Tensor = None
|
|
full_compress_sin: torch.Tensor = None
|
|
full_compress_cos: torch.Tensor = None
|
|
start_pos: torch.Tensor = None
|
|
num_reqs_actual: int | None = None
|
|
sas_metadata: torch.Tensor = None
|
|
qli_metadata: torch.Tensor = None
|
|
cu_cmp_seqlen_list: torch.Tensor = None
|
|
attn_mask: torch.Tensor | None = None
|
|
|
|
|
|
@dataclass
|
|
class AscendDSAMetadata:
|
|
"""Metadata for MLACommon.
|
|
NOTE: Please read the comment at the top of the file before trying to
|
|
understand this class
|
|
"""
|
|
|
|
num_actual_tokens: int # Number of tokens excluding padding.
|
|
query_start_loc: torch.Tensor
|
|
seq_lens: torch.Tensor
|
|
block_tables: torch.Tensor
|
|
sin: torch.Tensor
|
|
cos: torch.Tensor
|
|
|
|
num_decodes: int
|
|
num_decode_tokens: int
|
|
num_prefills: int
|
|
|
|
# For logging.
|
|
num_input_tokens: int = 0 # Number of tokens including padding.
|
|
|
|
# The dimension of the attention heads
|
|
head_dim: int | None = None
|
|
attn_mask: torch.Tensor = None
|
|
# chunked prefill by default if no attn_states passed
|
|
attn_state: AscendAttentionState = AscendAttentionState.ChunkedPrefill
|
|
|
|
req_metadata: AscendDSAReqMetadata | None = None
|
|
reshape_cache_event: torch.npu.Event = None
|
|
|
|
# metadata for dsv4 indexer
|
|
|
|
hadamard: torch.Tensor | None = None
|
|
|
|
start_pos: torch.Tensor | None = None
|
|
|
|
|
|
M = TypeVar("M", bound=AscendDSAMetadata)
|
|
|
|
|
|
class AscendDSACPMetadataBuilder(AttentionMetadataBuilder[AscendDSAMetadata]):
|
|
# Does this backend/builder support ACL Graphs for attention (default: no).
|
|
aclgraph_support: ClassVar[AttentionCGSupport] = AttentionCGSupport.UNIFORM_BATCH
|
|
hadamard = None
|
|
start_pos_prefill: torch.Tensor | None = None
|
|
req_sas_metadata: torch.Tensor
|
|
req_qli_metadata: torch.Tensor
|
|
block_size: int = 128
|
|
"""
|
|
NOTE: Please read the comment at the top of the file before trying to
|
|
understand this class
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
kv_cache_spec: AscendMLAAttentionSpec,
|
|
layer_names: list[str],
|
|
vllm_config: VllmConfig,
|
|
device: torch.device,
|
|
metadata_cls: type[AscendDSAMetadata] | None = None,
|
|
supports_dcp_with_varlen: bool = False,
|
|
):
|
|
self.kv_cache_spec = kv_cache_spec
|
|
self.metadata_cls = metadata_cls if metadata_cls is not None else AscendDSAMetadata
|
|
self.vllm_config = vllm_config
|
|
self.model_config = vllm_config.model_config
|
|
self.device = device
|
|
scheduler_config = vllm_config.scheduler_config
|
|
|
|
self.rope_dim = self.model_config.hf_text_config.qk_rope_head_dim
|
|
|
|
self.num_decodes = 0
|
|
self.num_prefills = 0
|
|
self.num_decode_tokens = 0
|
|
self.num_prefill_tokens = 0
|
|
self.num_actual_tokens: int | None = None
|
|
self.block_table: torch.Tensor = None
|
|
self.slot_mapping: torch.Tensor = None
|
|
self.seq_lens: torch.Tensor = None
|
|
self.seq_lens_cpu: torch.Tensor = None
|
|
|
|
self.compressor_ratio = getattr(kv_cache_spec, "compress_ratio", 0)
|
|
hf_config = self.model_config.hf_config
|
|
|
|
if AscendDSACPMetadataBuilder.hadamard is None:
|
|
if hf_config.model_type == "deepseek_v4":
|
|
indexer_head_dim = hf_config.index_head_dim
|
|
try:
|
|
from scipy.linalg import hadamard # type: ignore[import-untyped]
|
|
except ImportError as e:
|
|
raise ImportError(
|
|
"DeepSeek-V4 indexer attention requires SciPy for Hadamard transform. Please install scipy."
|
|
) from e
|
|
log_dim = math.ceil(math.log2(indexer_head_dim))
|
|
dim_padded = 2**log_dim
|
|
if self.vllm_config.model_config.enable_sleep_mode:
|
|
# Sleep mode allocates KV inside CaMemAllocator; tag Hadamard so
|
|
# sleep/wake does not treat it as KV cache.
|
|
from vllm_ascend.device_allocator.camem import CaMemAllocator
|
|
|
|
allocator = CaMemAllocator.get_instance()
|
|
with allocator.use_allocation_tag(CaMemAllocator.sleep_persistent_tag):
|
|
AscendDSACPMetadataBuilder.hadamard = torch.tensor(
|
|
hadamard(dim_padded, dtype=float), dtype=torch.float, device=self.device
|
|
).to(torch.bfloat16)
|
|
else:
|
|
AscendDSACPMetadataBuilder.hadamard = torch.tensor(
|
|
hadamard(dim_padded, dtype=float), dtype=torch.float, device=self.device
|
|
).to(torch.bfloat16)
|
|
self.start_pos_prefill = torch.zeros(scheduler_config.max_num_seqs, dtype=torch.int32, device=self.device)
|
|
self.req_sas_metadata = torch.zeros(1024, dtype=torch.int32, device=self.device)
|
|
self.req_qli_metadata = torch.zeros(1024, dtype=torch.int32, device=self.device)
|
|
self.cu_seqlens_ori_kv = torch.tensor([], device=self.device)
|
|
self.cu_seqlens_cmp_kv = torch.tensor([], device=self.device)
|
|
self.seqused_q = torch.tensor([], device=self.device)
|
|
self._zero_i32 = torch.tensor([0], device=self.device, dtype=torch.int32)
|
|
self.local_query_start_loc = torch.zeros(
|
|
scheduler_config.max_num_seqs + 1, dtype=torch.int32, device=self.device
|
|
)
|
|
self.local_seq_lens = torch.zeros(scheduler_config.max_num_seqs, dtype=torch.int32, device=self.device)
|
|
|
|
self.speculative_config = vllm_config.speculative_config
|
|
self.decode_threshold = 1
|
|
self.spec_slot_mapping = None
|
|
if get_ascend_device_type() in {AscendDeviceType.A5}:
|
|
self.slot_mapping_shape = (vllm_config.scheduler_config.max_num_batched_tokens,) # type: ignore
|
|
else:
|
|
self.slot_mapping_shape = (vllm_config.scheduler_config.max_num_batched_tokens, 2) # type: ignore
|
|
if self.speculative_config:
|
|
spec_token_num = self.speculative_config.num_speculative_tokens
|
|
self.spec_slot_mapping = [
|
|
torch.zeros(self.slot_mapping_shape, dtype=torch.int32, device=self.device)
|
|
for _ in range(spec_token_num)
|
|
]
|
|
self.spec_local_query_start_loc = [
|
|
torch.zeros(scheduler_config.max_num_seqs + 1, dtype=torch.int32, device=self.device)
|
|
for _ in range(spec_token_num)
|
|
]
|
|
self.spec_local_seq_lens = [
|
|
torch.zeros(scheduler_config.max_num_seqs, dtype=torch.int32, device=self.device)
|
|
for _ in range(spec_token_num)
|
|
]
|
|
self.decode_threshold += spec_token_num
|
|
assert self.decode_threshold <= 16, (
|
|
f"decode_threshold exceeded \
|
|
npu_fused_infer_attention_score TND layout's limit of 16, \
|
|
got {self.decode_threshold}"
|
|
)
|
|
|
|
self.reorder_batch_threshold = self.decode_threshold
|
|
# Note(qcs): we use two dimension slot_mapping for kvcache with shape
|
|
# [block_nums, block_size, head_num, head_dim]
|
|
self.slot_mapping = torch.zeros(self.slot_mapping_shape, dtype=torch.int32, device=self.device)
|
|
|
|
@classmethod
|
|
def get_cudagraph_support(
|
|
cls: type["AscendDSACPMetadataBuilder"],
|
|
vllm_config: VllmConfig,
|
|
kv_cache_spec: AttentionSpec,
|
|
) -> AttentionCGSupport:
|
|
# Explicit override in case the underlying builder specialized this getter.
|
|
# @override omitted only because of mypy limitation due to type variable.
|
|
return AttentionCGSupport.UNIFORM_BATCH
|
|
|
|
def build(
|
|
self,
|
|
common_prefix_len: int,
|
|
common_attn_metadata: AscendCommonAttentionMetadata,
|
|
fast_build: bool = False,
|
|
**kwargs,
|
|
) -> AscendDSAMetadata:
|
|
num_reqs = common_attn_metadata.num_reqs
|
|
query_start_loc = common_attn_metadata.query_start_loc
|
|
num_reqs_actual = kwargs.get("num_reqs_actual")
|
|
self.block_size = kwargs.get("block_size", 128)
|
|
|
|
common_ratio_to_sas_metadata = kwargs.get("common_ratio_to_sas_metadata")
|
|
assert common_ratio_to_sas_metadata is not None
|
|
self.common_ratio_to_sas_metadata = common_ratio_to_sas_metadata
|
|
self.num_actual_tokens = common_attn_metadata.num_actual_tokens
|
|
attn_state = kwargs.get("attn_state", common_attn_metadata.attn_state)
|
|
has_prefill = _has_prefill(attn_state)
|
|
|
|
num_input_tokens = common_attn_metadata.num_input_tokens
|
|
if self.common_ratio_to_sas_metadata.get("input_positions", None) is None:
|
|
self.num_decodes, self.num_prefills, self.num_decode_tokens, self.num_prefill_tokens = (
|
|
split_decodes_and_prefills(
|
|
common_attn_metadata,
|
|
decode_threshold=self.decode_threshold,
|
|
treat_short_extends_as_decodes=False,
|
|
)
|
|
)
|
|
self.common_ratio_to_sas_metadata["num_decodes"] = self.num_decodes
|
|
self.common_ratio_to_sas_metadata["num_prefills"] = self.num_prefills
|
|
self.common_ratio_to_sas_metadata["num_decode_tokens"] = self.num_decode_tokens
|
|
self.common_ratio_to_sas_metadata["num_prefill_tokens"] = self.num_prefill_tokens
|
|
input_positions = common_attn_metadata.positions[:num_input_tokens].long()
|
|
input_positions_cpu = common_attn_metadata.positions_cpu[:num_input_tokens].long()
|
|
self.common_ratio_to_sas_metadata["input_positions"] = input_positions
|
|
self.common_ratio_to_sas_metadata["input_positions_cpu"] = input_positions_cpu
|
|
cos, sin = get_cos_and_sin_dsa(input_positions, use_cache=not has_prefill)
|
|
self.common_ratio_to_sas_metadata["cos"] = cos
|
|
self.common_ratio_to_sas_metadata["sin"] = sin
|
|
self.seq_lens = common_attn_metadata.seq_lens[:num_reqs]
|
|
self.common_ratio_to_sas_metadata["seq_lens"] = self.seq_lens
|
|
# Prefer _seq_lens_cpu (always available, updated during draft
|
|
# iterations) over seq_lens_cpu (None in async spec decode mode).
|
|
if common_attn_metadata._seq_lens_cpu is not None:
|
|
_seq_lens_cpu = common_attn_metadata._seq_lens_cpu
|
|
elif common_attn_metadata.seq_lens_cpu is not None:
|
|
_seq_lens_cpu = common_attn_metadata.seq_lens_cpu
|
|
else:
|
|
_seq_lens_cpu = common_attn_metadata.seq_lens.cpu()
|
|
self.seq_lens_cpu = _seq_lens_cpu
|
|
self.common_ratio_to_sas_metadata["seq_lens_cpu"] = self.seq_lens_cpu
|
|
else:
|
|
self.num_decodes, self.num_prefills, self.num_decode_tokens, self.num_prefill_tokens = (
|
|
self.common_ratio_to_sas_metadata["num_decodes"],
|
|
self.common_ratio_to_sas_metadata["num_prefills"],
|
|
self.common_ratio_to_sas_metadata["num_decode_tokens"],
|
|
self.common_ratio_to_sas_metadata["num_prefill_tokens"],
|
|
)
|
|
input_positions = self.common_ratio_to_sas_metadata["input_positions"]
|
|
input_positions_cpu = self.common_ratio_to_sas_metadata["input_positions_cpu"]
|
|
cos, sin = self.common_ratio_to_sas_metadata["cos"], self.common_ratio_to_sas_metadata["sin"]
|
|
self.seq_lens = self.common_ratio_to_sas_metadata["seq_lens"]
|
|
self.seq_lens_cpu = self.common_ratio_to_sas_metadata["seq_lens_cpu"]
|
|
|
|
slot_mapping = common_attn_metadata.slot_mapping[:num_input_tokens]
|
|
self.slot_mapping[:num_input_tokens] = DeviceOperator.format_dsa_slot_mapping(slot_mapping, self.block_size)
|
|
|
|
self.block_table = common_attn_metadata.block_table_tensor[:num_reqs]
|
|
|
|
req_metadata = self.build_req_metadata(
|
|
common_attn_metadata, input_positions, input_positions_cpu, num_input_tokens, num_reqs_actual, attn_state
|
|
)
|
|
|
|
return self.metadata_cls( # type: ignore
|
|
num_input_tokens=common_attn_metadata.num_input_tokens,
|
|
num_actual_tokens=self.num_actual_tokens,
|
|
head_dim=self.model_config.get_head_size(),
|
|
attn_mask=None,
|
|
num_decodes=self.num_decodes,
|
|
num_decode_tokens=self.num_decode_tokens,
|
|
num_prefills=self.num_prefills,
|
|
attn_state=attn_state,
|
|
req_metadata=req_metadata,
|
|
query_start_loc=query_start_loc,
|
|
block_tables=None,
|
|
seq_lens=self.seq_lens,
|
|
cos=cos,
|
|
sin=sin,
|
|
hadamard=AscendDSACPMetadataBuilder.hadamard,
|
|
)
|
|
|
|
def build_for_drafting(
|
|
self,
|
|
common_attn_metadata: AscendCommonAttentionMetadata,
|
|
draft_index: int,
|
|
fast_build: bool = False,
|
|
**kwargs,
|
|
) -> AscendDSAMetadata:
|
|
assert self.compressor_ratio <= 1, "vLLM-Ascend only support SWA-layer for Deepseek-V4 now."
|
|
num_reqs = common_attn_metadata.num_reqs
|
|
num_input_tokens = common_attn_metadata.num_input_tokens
|
|
num_decodes, num_prefills, num_decode_tokens, _ = split_decodes_and_prefills(
|
|
common_attn_metadata,
|
|
decode_threshold=self.decode_threshold,
|
|
treat_short_extends_as_decodes=False,
|
|
)
|
|
|
|
self.num_decodes = num_decodes
|
|
self.num_prefills = num_prefills
|
|
self.num_decode_tokens = num_decode_tokens
|
|
self.num_actual_tokens = common_attn_metadata.num_actual_tokens
|
|
self.seq_lens = common_attn_metadata.seq_lens[:num_reqs]
|
|
self.block_size = kwargs.get("block_size", 128)
|
|
|
|
input_positions = common_attn_metadata.positions[:num_input_tokens].long()
|
|
# Draft steps update positions independently. Reusing the global RoPE
|
|
# cache can let later draft steps overwrite step-0 metadata.
|
|
cos, sin = get_cos_and_sin_dsa(input_positions, use_cache=False)
|
|
|
|
slot_mapping = common_attn_metadata.slot_mapping[:num_input_tokens]
|
|
|
|
assert self.spec_slot_mapping is not None
|
|
self.spec_slot_mapping[draft_index - 1][:num_input_tokens] = DeviceOperator.format_dsa_slot_mapping(
|
|
slot_mapping, self.block_size
|
|
)
|
|
|
|
self.block_table = common_attn_metadata.block_table_tensor[:num_reqs]
|
|
req_metadata = self.build_req_metadata_for_drafting(
|
|
draft_index=draft_index,
|
|
common_attn_metadata=common_attn_metadata,
|
|
input_positions=input_positions,
|
|
num_input_tokens=num_input_tokens,
|
|
)
|
|
|
|
return self.metadata_cls( # type: ignore
|
|
num_input_tokens=common_attn_metadata.num_input_tokens,
|
|
num_actual_tokens=self.num_actual_tokens,
|
|
head_dim=self.model_config.get_head_size(),
|
|
attn_mask=None,
|
|
num_decodes=num_decodes,
|
|
num_decode_tokens=num_decode_tokens,
|
|
num_prefills=num_prefills,
|
|
attn_state=common_attn_metadata.attn_state,
|
|
req_metadata=req_metadata,
|
|
query_start_loc=common_attn_metadata.query_start_loc,
|
|
block_tables=None,
|
|
seq_lens=self.seq_lens,
|
|
cos=cos,
|
|
sin=sin,
|
|
hadamard=None,
|
|
)
|
|
|
|
def build_req_metadata_for_drafting(
|
|
self,
|
|
draft_index: int,
|
|
common_attn_metadata: AscendCommonAttentionMetadata,
|
|
input_positions: torch.Tensor,
|
|
num_input_tokens: int,
|
|
) -> AscendDSAReqMetadata:
|
|
"""Build DSA-CP metadata for one draft step."""
|
|
num_reqs = common_attn_metadata.num_reqs
|
|
query_start_loc = common_attn_metadata.query_start_loc
|
|
query_start_loc_cpu = common_attn_metadata.query_start_loc_cpu
|
|
seq_lens_q = query_start_loc[1:] - query_start_loc[:-1]
|
|
has_prefill = _has_prefill(common_attn_metadata.attn_state)
|
|
|
|
cos, sin = get_cos_and_sin_dsa(input_positions, use_cache=False)
|
|
(
|
|
local_start,
|
|
local_end_with_pad,
|
|
tokens_per_rank,
|
|
num_tokens_pad,
|
|
local_query_start_loc,
|
|
local_seq_lens,
|
|
local_cos,
|
|
local_sin,
|
|
) = self._build_local_token_metadata(
|
|
num_reqs=num_reqs,
|
|
num_input_tokens=num_input_tokens,
|
|
input_positions=input_positions,
|
|
query_start_loc=query_start_loc,
|
|
seq_lens=self.seq_lens[:num_reqs],
|
|
use_cache=False,
|
|
local_query_start_loc=self.spec_local_query_start_loc[draft_index - 1],
|
|
local_seq_lens=self.spec_local_seq_lens[draft_index - 1],
|
|
)
|
|
local_query_start_loc = local_query_start_loc.clone()
|
|
local_seq_lens = local_seq_lens.clone()
|
|
|
|
_, _, _, _, local_query_start_loc_cpu, local_seq_lens_cpu, _, _ = self._build_local_token_metadata(
|
|
num_reqs=num_reqs,
|
|
num_input_tokens=num_input_tokens,
|
|
input_positions=None,
|
|
query_start_loc=query_start_loc_cpu,
|
|
seq_lens=self.seq_lens_cpu[:num_reqs],
|
|
use_cache=False,
|
|
)
|
|
local_seq_lens_q_cpu = local_query_start_loc_cpu[1 : num_reqs + 1] - local_query_start_loc_cpu[:num_reqs]
|
|
max_local_query_len = max(1, int(local_seq_lens_q_cpu.max().item()))
|
|
max_local_seq_lens = max(1, int(local_seq_lens_cpu.max().item()))
|
|
|
|
start_pos = self.seq_lens[:num_reqs] - seq_lens_q
|
|
|
|
assert self.spec_slot_mapping is not None
|
|
slot_mapping = self.spec_slot_mapping[draft_index - 1][: self.num_actual_tokens]
|
|
|
|
num_heads = self.model_config.hf_config.num_attention_heads
|
|
metadata_op = DeviceOperator.get_dsa_sparse_attn_metadata_op()
|
|
metadata_kwargs = DeviceOperator.get_dsa_sparse_attn_metadata_kwargs(self.seqused_q.device)
|
|
metadata_kwargs.setdefault("device", str(self.seqused_q.device))
|
|
cu_seqlens_ori_kv = (
|
|
local_query_start_loc
|
|
if has_prefill
|
|
else DeviceOperator.get_dsa_decode_cu_seqlens_ori_kv(
|
|
None,
|
|
"draft_cu_seqlens_ori_kv",
|
|
local_seq_lens,
|
|
num_reqs,
|
|
self._zero_i32,
|
|
self.cu_seqlens_ori_kv,
|
|
)
|
|
)
|
|
cu_seqlens_cmp_kv = (
|
|
None if has_prefill else DeviceOperator.get_dsa_decode_cu_seqlens_cmp_kv(self.cu_seqlens_cmp_kv)
|
|
)
|
|
sas_metadata = metadata_op(
|
|
**metadata_kwargs,
|
|
num_heads_q=num_heads,
|
|
num_heads_kv=1,
|
|
head_dim=self.model_config.get_head_size(),
|
|
cu_seqlens_q=local_query_start_loc,
|
|
cu_seqlens_ori_kv=cu_seqlens_ori_kv,
|
|
cu_seqlens_cmp_kv=cu_seqlens_cmp_kv,
|
|
seqused_q=self.seqused_q,
|
|
seqused_kv=local_seq_lens,
|
|
max_seqlen_q=max_local_query_len,
|
|
max_seqlen_kv=max_local_seq_lens,
|
|
batch_size=num_reqs,
|
|
cmp_ratio=1,
|
|
ori_mask_mode=4,
|
|
ori_win_left=self.model_config.hf_config.sliding_window - 1,
|
|
ori_win_right=0,
|
|
layout_q="TND",
|
|
layout_kv="PA_ND",
|
|
has_ori_kv=True,
|
|
has_cmp_kv=False,
|
|
)
|
|
|
|
cp_metadata = DSACPMetadata(
|
|
local_query_start_loc=local_query_start_loc,
|
|
local_seq_lens=local_seq_lens,
|
|
local_start=local_start,
|
|
local_end=local_end_with_pad,
|
|
tokens_per_rank=tokens_per_rank,
|
|
num_tokens_pad=num_tokens_pad,
|
|
local_sin=local_sin,
|
|
local_cos=local_cos,
|
|
)
|
|
|
|
return AscendDSAReqMetadata(
|
|
input_positions=input_positions,
|
|
block_table=self.block_table[:num_reqs, ...],
|
|
slot_mapping=slot_mapping,
|
|
block_size=self.block_size,
|
|
seq_lens=self.seq_lens[:num_reqs],
|
|
query_start_loc=query_start_loc,
|
|
cp_metadata=cp_metadata,
|
|
sin=sin,
|
|
cos=cos,
|
|
start_pos=start_pos,
|
|
sas_metadata=sas_metadata,
|
|
qli_metadata=None,
|
|
cu_cmp_seqlen_list=None,
|
|
)
|
|
|
|
def _num_compressor_metadata_rows(
|
|
self,
|
|
common_attn_metadata: AscendCommonAttentionMetadata,
|
|
) -> int:
|
|
assert self.num_actual_tokens is not None
|
|
num_tokens = self.num_actual_tokens
|
|
return min(num_tokens, num_tokens // self.compressor_ratio + common_attn_metadata.num_reqs)
|
|
|
|
def build_req_metadata(
|
|
self,
|
|
common_attn_metadata: AscendCommonAttentionMetadata,
|
|
input_positions: torch.Tensor,
|
|
input_positions_cpu: torch.Tensor,
|
|
num_input_tokens: int,
|
|
num_reqs_actual: int | None,
|
|
attn_state: AscendAttentionState,
|
|
) -> AscendDSAReqMetadata:
|
|
"""Build a single unified metadata for all requests (prefill + decode)."""
|
|
num_reqs = common_attn_metadata.num_reqs
|
|
has_prefill = _has_prefill(attn_state)
|
|
query_start_loc = common_attn_metadata.query_start_loc
|
|
query_start_loc_cpu = common_attn_metadata.query_start_loc_cpu
|
|
|
|
seq_lens_q = query_start_loc[1:] - query_start_loc[:-1]
|
|
|
|
# cos/sin for all tokens
|
|
cos, sin = get_cos_and_sin_dsa(input_positions, use_cache=not has_prefill)
|
|
|
|
(
|
|
local_start,
|
|
local_end_with_pad,
|
|
tokens_per_rank,
|
|
num_tokens_pad,
|
|
local_query_start_loc,
|
|
local_seq_lens,
|
|
local_cos,
|
|
local_sin,
|
|
) = self._build_local_token_metadata(
|
|
num_reqs=num_reqs,
|
|
num_input_tokens=num_input_tokens,
|
|
input_positions=input_positions,
|
|
query_start_loc=query_start_loc,
|
|
seq_lens=self.seq_lens[:num_reqs],
|
|
use_cache=not has_prefill,
|
|
local_query_start_loc=self.local_query_start_loc,
|
|
local_seq_lens=self.local_seq_lens,
|
|
)
|
|
local_seq_lens_q = local_query_start_loc[1 : num_reqs + 1] - local_query_start_loc[:num_reqs]
|
|
|
|
_, _, _, _, local_query_start_loc_cpu, local_seq_lens_cpu, _, _ = self._build_local_token_metadata(
|
|
num_reqs=num_reqs,
|
|
num_input_tokens=num_input_tokens,
|
|
input_positions=None,
|
|
query_start_loc=query_start_loc_cpu,
|
|
seq_lens=self.seq_lens_cpu[:num_reqs],
|
|
use_cache=False,
|
|
)
|
|
local_seq_lens_q_cpu = local_query_start_loc_cpu[1 : num_reqs + 1] - local_query_start_loc_cpu[:num_reqs]
|
|
max_local_query_len = max(1, int(local_seq_lens_q_cpu.max().item()))
|
|
max_local_seq_lens = max(1, int(local_seq_lens_cpu.max().item()))
|
|
|
|
# start_pos: context length before current query
|
|
start_pos = self.seq_lens[:num_reqs] - seq_lens_q
|
|
|
|
assert self.start_pos_prefill is not None
|
|
self.start_pos_prefill.fill_(0)
|
|
self.start_pos_prefill[:num_reqs] = start_pos
|
|
|
|
if num_reqs_actual is None:
|
|
num_reqs_actual = num_reqs
|
|
else:
|
|
num_reqs_actual = min(num_reqs_actual, num_reqs)
|
|
if num_reqs_actual < num_reqs:
|
|
self.start_pos_prefill[num_reqs_actual:].fill_(0)
|
|
self.block_table[num_reqs_actual:num_reqs, ...].fill_(0)
|
|
|
|
# --- Compressed positions ---
|
|
full_compress_cos, full_compress_sin = None, None
|
|
cu_cmp_seqlens = self._get_cmp_seqlens_for_metadata(has_prefill)
|
|
|
|
if self.compressor_ratio > 1:
|
|
layer_name = f"c{self.compressor_ratio}"
|
|
# Keep only graph inputs here. The compressor metadata op itself is
|
|
# launched in forward at the real compressor consumer.
|
|
num_compressed_tokens = self._num_compressor_metadata_rows(common_attn_metadata)
|
|
full_compress_cos, full_compress_sin = get_full_cos_and_sin_dsa(layer_name)
|
|
slot_mapping = None
|
|
else:
|
|
num_compressed_tokens = None
|
|
slot_mapping = self.slot_mapping[: self.num_actual_tokens]
|
|
|
|
# --- SAS metadata (all requests combined) ---
|
|
num_heads = self.model_config.hf_config.num_attention_heads
|
|
index_topk = self.model_config.hf_config.index_topk
|
|
|
|
sas_metadata = self._build_sas_metadata(
|
|
num_heads=num_heads,
|
|
query_start_loc=local_query_start_loc,
|
|
seq_lens=local_seq_lens,
|
|
seq_lens_q=local_seq_lens_q,
|
|
max_query_len=max_local_query_len,
|
|
max_seq_lens=max_local_seq_lens,
|
|
index_topk=index_topk,
|
|
num_reqs=num_reqs,
|
|
has_prefill=has_prefill,
|
|
cu_cmp_seqlen_list=cu_cmp_seqlens,
|
|
)
|
|
|
|
# --- QLI metadata (all requests combined) ---
|
|
qli_metadata = self._build_qli_metadata(
|
|
query_start_loc=local_query_start_loc,
|
|
seq_lens=local_seq_lens,
|
|
seq_lens_q=local_seq_lens_q,
|
|
num_reqs=num_reqs,
|
|
)
|
|
|
|
cp_metadata = DSACPMetadata(
|
|
local_query_start_loc=local_query_start_loc,
|
|
local_seq_lens=local_seq_lens,
|
|
local_start=local_start,
|
|
local_end=local_end_with_pad,
|
|
tokens_per_rank=tokens_per_rank,
|
|
num_tokens_pad=num_tokens_pad,
|
|
local_sin=local_sin,
|
|
local_cos=local_cos,
|
|
)
|
|
|
|
return AscendDSAReqMetadata(
|
|
input_positions=input_positions,
|
|
block_table=self.block_table[:num_reqs, ...],
|
|
slot_mapping=slot_mapping,
|
|
block_size=self.block_size,
|
|
seq_lens=self.seq_lens[:num_reqs],
|
|
query_start_loc=query_start_loc,
|
|
cp_metadata=cp_metadata,
|
|
sin=sin,
|
|
cos=cos,
|
|
full_compress_sin=full_compress_sin,
|
|
full_compress_cos=full_compress_cos,
|
|
start_pos=self.start_pos_prefill[:num_reqs],
|
|
num_compressed_tokens=num_compressed_tokens,
|
|
num_reqs_actual=num_reqs_actual,
|
|
sas_metadata=sas_metadata,
|
|
qli_metadata=qli_metadata,
|
|
cu_cmp_seqlen_list=cu_cmp_seqlens,
|
|
)
|
|
|
|
def _build_local_token_metadata(
|
|
self,
|
|
num_reqs,
|
|
num_input_tokens,
|
|
input_positions,
|
|
query_start_loc,
|
|
seq_lens,
|
|
use_cache,
|
|
local_query_start_loc=None,
|
|
local_seq_lens=None,
|
|
):
|
|
"""
|
|
For example:
|
|
If we have TP size 3, num_input_tokens=45, and
|
|
query_start_loc = [0, 1, 3, 6, 10, 15, 21, 28, 36, 45].
|
|
That means we have 9 requests with seq lens [1, 2, 3, 4, 5, 6, 7, 8, 9].
|
|
For tp_rank 1, local_start=15, local_end=30, tokens_per_rank=15.
|
|
local_query_start=[15, 15, 15, 15, 15, 15, 21, 28, 30]
|
|
local_query_end = [15, 15, 15, 15, 15, 21, 28, 30, 30]
|
|
local_query_lens = [0, 0, 0, 0, 0, 6, 7, 2, 0]
|
|
self.local_query_start_loc = [0, 0, 0, 0, 0, 0, 6, 13, 15]
|
|
offset = [-14, -12, -9, -5, 0, 0, 0, 6, 15]
|
|
seq_lens-offset=[15, 14, 12, 9, 5, 6, 7, 2, -6]
|
|
local_reqs_mask = [0, 0, 0, 0, 0, 1, 1, 1, 0]
|
|
local_seq_lens = [0, 0, 0, 0, 0, 6, 7, 2, 0]
|
|
"""
|
|
tp_group = get_tp_group()
|
|
tp_size = tp_group.world_size
|
|
tp_rank = tp_group.rank_in_group
|
|
# Split the flattened token stream evenly across TP ranks. Padding keeps
|
|
# every rank's local slice the same length, which simplifies CP kernels.
|
|
num_tokens_pad = ((num_input_tokens + tp_size - 1) // tp_size) * tp_size
|
|
tokens_per_rank = num_tokens_pad // tp_size
|
|
local_start = tp_rank * tokens_per_rank
|
|
local_end = local_start + tokens_per_rank
|
|
|
|
if local_query_start_loc is not None:
|
|
local_query_start_loc.fill_(0)
|
|
local_seq_lens.fill_(0)
|
|
|
|
# Intersect each request's global token interval with this rank's local
|
|
# token interval, then build the per-rank query_start_loc from lengths.
|
|
local_query_start = torch.clamp(query_start_loc[:-1], min=local_start, max=local_end)
|
|
local_query_end = torch.clamp(query_start_loc[1:], min=local_start, max=local_end)
|
|
local_query_lens = local_query_end - local_query_start
|
|
if local_query_start_loc is not None:
|
|
local_query_start_loc[1 : num_reqs + 1] = torch.cumsum(local_query_lens, dim=0)
|
|
else:
|
|
local_query_start_loc = torch.cat(
|
|
[
|
|
torch.tensor([0], dtype=local_query_lens.dtype, device=local_query_lens.device),
|
|
torch.cumsum(local_query_lens, dim=0),
|
|
],
|
|
0,
|
|
)
|
|
|
|
# For requests that cross the local slice boundary, offset removes the
|
|
# tokens that live on later ranks so local_seq_lens matches local queries.
|
|
offset = query_start_loc[1:] - local_query_end
|
|
valid_local_req = (local_query_lens > 0) & (seq_lens > 0)
|
|
safe_local_seq_lens = torch.clamp_min(seq_lens - offset, 0)
|
|
safe_local_seq_lens = torch.where(
|
|
valid_local_req,
|
|
safe_local_seq_lens,
|
|
torch.zeros_like(safe_local_seq_lens),
|
|
)
|
|
if local_seq_lens is not None:
|
|
local_seq_lens[:num_reqs] = safe_local_seq_lens
|
|
else:
|
|
local_seq_lens = safe_local_seq_lens
|
|
|
|
# RoPE tables are generated on the padded global positions first, then
|
|
# sliced to this rank so local tokens keep their original positions.
|
|
if input_positions is not None:
|
|
pad_tokens = num_tokens_pad - input_positions.shape[0]
|
|
if pad_tokens > 0:
|
|
input_positions = F.pad(input_positions, (0, pad_tokens), value=0)
|
|
local_cos, local_sin = get_cos_and_sin_dsa(input_positions, use_cache=use_cache)
|
|
local_cos = local_cos[local_start:local_end]
|
|
local_sin = local_sin[local_start:local_end]
|
|
else:
|
|
local_cos = None
|
|
local_sin = None
|
|
return (
|
|
local_start,
|
|
local_end,
|
|
tokens_per_rank,
|
|
num_tokens_pad,
|
|
local_query_start_loc[: num_reqs + 1],
|
|
local_seq_lens[:num_reqs],
|
|
local_cos,
|
|
local_sin,
|
|
)
|
|
|
|
def _get_cmp_seqlens_for_metadata(self, has_prefill):
|
|
if self.compressor_ratio <= 1:
|
|
return None
|
|
if has_prefill:
|
|
return None
|
|
return DeviceOperator.get_dsa_decode_cu_seqlens_cmp_kv(self.cu_seqlens_cmp_kv)
|
|
|
|
def _build_sas_metadata(
|
|
self,
|
|
num_heads,
|
|
query_start_loc,
|
|
seq_lens,
|
|
seq_lens_q,
|
|
max_query_len,
|
|
max_seq_lens,
|
|
index_topk,
|
|
num_reqs,
|
|
has_prefill,
|
|
cu_cmp_seqlen_list,
|
|
):
|
|
cmp_ratio = self.compressor_ratio if self.compressor_ratio > 1 else 1
|
|
cache_key = f"cp_sas_c{cmp_ratio}"
|
|
metadata = self.common_ratio_to_sas_metadata.get(cache_key)
|
|
if metadata is None:
|
|
cu_seqlens_ori_kv = (
|
|
query_start_loc
|
|
if has_prefill
|
|
else DeviceOperator.get_dsa_decode_cu_seqlens_ori_kv(
|
|
self.common_ratio_to_sas_metadata,
|
|
f"{cache_key}_cu_seqlens_ori_kv",
|
|
seq_lens,
|
|
num_reqs,
|
|
self._zero_i32,
|
|
self.cu_seqlens_ori_kv,
|
|
)
|
|
)
|
|
cu_seqlens_cmp_kv = (
|
|
None if has_prefill else DeviceOperator.get_dsa_decode_cu_seqlens_cmp_kv(self.cu_seqlens_cmp_kv)
|
|
)
|
|
metadata_op = DeviceOperator.get_dsa_sparse_attn_metadata_op()
|
|
metadata_kwargs = DeviceOperator.get_dsa_sparse_attn_metadata_kwargs(self.seqused_q.device)
|
|
metadata_kwargs.setdefault("device", str(self.seqused_q.device))
|
|
kw = dict(
|
|
**metadata_kwargs,
|
|
num_heads_q=num_heads,
|
|
num_heads_kv=1,
|
|
head_dim=self.model_config.get_head_size(),
|
|
cu_seqlens_q=query_start_loc,
|
|
cu_seqlens_ori_kv=cu_seqlens_ori_kv,
|
|
cu_seqlens_cmp_kv=cu_seqlens_cmp_kv,
|
|
seqused_q=self.seqused_q,
|
|
seqused_kv=seq_lens,
|
|
max_seqlen_q=max_query_len,
|
|
max_seqlen_kv=max_seq_lens,
|
|
batch_size=num_reqs,
|
|
ori_mask_mode=4,
|
|
ori_win_left=self.model_config.hf_config.sliding_window - 1,
|
|
ori_win_right=0,
|
|
layout_q="TND",
|
|
layout_kv="PA_ND",
|
|
has_ori_kv=True,
|
|
)
|
|
|
|
if self.compressor_ratio > 1:
|
|
kw["has_cmp_kv"] = True
|
|
if self.compressor_ratio == 4:
|
|
kw["cmp_mask_mode"] = 3
|
|
kw["cmp_topk"] = index_topk
|
|
else:
|
|
kw["cmp_mask_mode"] = 3
|
|
kw["cmp_ratio"] = cmp_ratio
|
|
kw["cu_seqlens_cmp_kv"] = cu_cmp_seqlen_list
|
|
else:
|
|
kw["cmp_ratio"] = cmp_ratio
|
|
kw["has_cmp_kv"] = False
|
|
|
|
metadata = metadata_op(**kw)
|
|
self.common_ratio_to_sas_metadata[cache_key] = metadata
|
|
self.req_sas_metadata[:1024] = metadata
|
|
return self.req_sas_metadata[:1024]
|
|
|
|
def _build_qli_metadata(self, query_start_loc, seq_lens, seq_lens_q, num_reqs):
|
|
if self.compressor_ratio != 4:
|
|
return None
|
|
|
|
cache_key = "cp_qli"
|
|
metadata = self.common_ratio_to_sas_metadata.get(cache_key)
|
|
|
|
if metadata is None:
|
|
max_seqlen_q = max(1, int(seq_lens_q.max().item()))
|
|
max_seqlen_k = max(1, int(seq_lens.max().item()))
|
|
metadata = torch.ops._C_ascend.npu_vllm_quant_lightning_indexer_metadata(
|
|
actual_seq_lengths_query=query_start_loc[1:].clone(),
|
|
actual_seq_lengths_key=seq_lens.clone(),
|
|
num_heads_q=self.model_config.hf_config.index_n_heads,
|
|
num_heads_k=1,
|
|
head_dim=self.model_config.hf_config.index_head_dim,
|
|
query_quant_mode=0,
|
|
key_quant_mode=0,
|
|
batch_size=num_reqs,
|
|
max_seqlen_q=max_seqlen_q,
|
|
max_seqlen_k=max_seqlen_k,
|
|
layout_query="TND",
|
|
layout_key="PA_BSND",
|
|
sparse_count=self.model_config.hf_config.index_topk,
|
|
sparse_mode=3,
|
|
pre_tokens=(1 << 63) - 1,
|
|
next_tokens=(1 << 63) - 1,
|
|
cmp_ratio=4,
|
|
device=str(self.seqused_q.device),
|
|
)
|
|
self.common_ratio_to_sas_metadata[cache_key] = metadata
|
|
self.req_qli_metadata[:1024] = metadata
|
|
return self.req_qli_metadata[:1024]
|
|
|
|
def build_for_graph_capture(
|
|
self,
|
|
common_attn_metadata: AscendCommonAttentionMetadata,
|
|
attn_state: AscendAttentionState = AscendAttentionState.DecodeOnly,
|
|
**kwargs,
|
|
):
|
|
if attn_state in {AscendAttentionState.DecodeOnly, AscendAttentionState.SpecDecoding}:
|
|
attn_metadata = self.build(
|
|
common_prefix_len=0,
|
|
common_attn_metadata=common_attn_metadata,
|
|
attn_state=attn_state,
|
|
**kwargs,
|
|
)
|
|
else:
|
|
raise NotImplementedError(
|
|
f"Graph capture only supports DecodeOnly and SpecDecoding attn states, got {attn_state}."
|
|
)
|
|
|
|
assert attn_metadata is not None
|
|
return attn_metadata
|
|
|
|
|
|
class AscendDSACPImpl(DSAAttentionImpl):
|
|
"""
|
|
NOTE: Please read the comment at the top of the file before trying to
|
|
understand this class
|
|
"""
|
|
|
|
wo_a_full_pool: ClassVar[torch.Tensor | None] = None
|
|
wo_a_full_weight_scale_pool: ClassVar[torch.Tensor | None] = None
|
|
wo_b_full_pool: ClassVar[torch.Tensor | None] = None
|
|
wo_b_full_weight_scale_pool: ClassVar[torch.Tensor | None] = None
|
|
|
|
def __init__(
|
|
self,
|
|
n_heads: int,
|
|
scale: float,
|
|
n_local_heads: int,
|
|
q_lora_rank: int,
|
|
o_lora_rank: int,
|
|
head_dim: int,
|
|
rope_head_dim: int | None,
|
|
nope_head_dim: int,
|
|
n_groups: int,
|
|
n_local_groups: int,
|
|
window_size: int,
|
|
compress_ratio: int,
|
|
**kwargs,
|
|
):
|
|
self.num_heads = n_heads
|
|
self.n_local_heads = n_local_heads
|
|
self.scale = scale
|
|
self.o_lora_rank = o_lora_rank
|
|
self.nope_head_dim = nope_head_dim
|
|
self.rope_head_dim = rope_head_dim
|
|
self.head_dim = head_dim
|
|
self.n_group = n_groups
|
|
self.n_local_groups = n_local_groups
|
|
self.window_size = window_size
|
|
self.q_lora_rank = q_lora_rank
|
|
self.compress_ratio = compress_ratio
|
|
self.softmax_scale = self.head_dim**-0.5
|
|
self.tp_group = get_tp_group()
|
|
self.tp_size = self.tp_group.world_size
|
|
self.tp_rank = self.tp_group.rank_in_group
|
|
|
|
# MLA Args
|
|
self.wq_a = kwargs["wq_a"]
|
|
self.wq_b = kwargs["wq_b"]
|
|
self.wkv = kwargs["wkv"]
|
|
self.q_norm = kwargs["q_norm"]
|
|
self.q_norm_without_weight = kwargs.get("q_norm_without_weight")
|
|
self.kv_norm = kwargs["kv_norm"]
|
|
|
|
self.indexer = kwargs.get("indexer")
|
|
self.compressor = kwargs.get("compressor")
|
|
|
|
self.wo_a = kwargs["wo_a"]
|
|
self.wo_b = kwargs["wo_b"]
|
|
|
|
self.enable_dsa_cp_with_o_proj_tp = enable_dsa_cp_with_o_proj_tp() and (
|
|
get_ascend_device_type() == AscendDeviceType.A5
|
|
)
|
|
self._wo_a_dynamic_quant = False
|
|
self._wo_b_dynamic_quant = False
|
|
|
|
self.eps = kwargs["eps"]
|
|
|
|
self.attn_sink = kwargs["attn_sink"]
|
|
|
|
self.vllm_config = get_current_vllm_config()
|
|
|
|
# indexer param
|
|
if self.indexer is not None:
|
|
self.indexer_heads: int = self.indexer.n_heads
|
|
self.inderxer_dim: int = self.indexer.head_dim
|
|
self.inderxer_wq_b = self.indexer.wq_b
|
|
self.weights_proj = self.indexer.weights_proj
|
|
self.indexer_softmax_scale = self.inderxer_dim**-0.5
|
|
|
|
self.indexer_compress = self.indexer.compressor
|
|
|
|
# indexer_compressor
|
|
self.indexcom_ape = self.indexer.compressor.ape
|
|
self.indexcom_wkv = self.indexer.compressor.wkv
|
|
self.indexcom_wgate = self.indexer.compressor.wgate
|
|
self.indexcom_norm = self.indexer.compressor.norm
|
|
|
|
self.indexcom_head_dim = self.indexer.compressor.head_dim
|
|
self.indexcom_rotate = self.indexer.compressor.rotate
|
|
self.index_topk = self.indexer.index_topk
|
|
|
|
# compress param
|
|
if self.compressor is not None:
|
|
self.compressor_head_dim = self.compressor.head_dim
|
|
self.compressor_overlap = self.compressor.overlap
|
|
self.compressor_rotate = self.compressor.rotate
|
|
|
|
self.compressor_ape = self.compressor.ape
|
|
self.compressor_wkv = self.compressor.wkv
|
|
self.compressor_wgate = self.compressor.wgate
|
|
self.compressor_norm = self.compressor.norm
|
|
self.compressor_norm_eps = self.compressor.norm_eps
|
|
|
|
def _compute_compressor_metadata(
|
|
self,
|
|
metadata: AscendDSAReqMetadata,
|
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
|
assert metadata.full_compress_cos is not None
|
|
assert metadata.full_compress_sin is not None
|
|
assert metadata.num_compressed_tokens is not None
|
|
assert metadata.start_pos is not None
|
|
assert metadata.num_reqs_actual is not None
|
|
full_compress_cos = metadata.full_compress_cos.view(
|
|
metadata.full_compress_cos.shape[0],
|
|
metadata.full_compress_cos.shape[-1],
|
|
)
|
|
full_compress_sin = metadata.full_compress_sin.view(
|
|
metadata.full_compress_sin.shape[0],
|
|
metadata.full_compress_sin.shape[-1],
|
|
)
|
|
return torch.ops._C_ascend.compressor_metadata(
|
|
full_compress_cos,
|
|
full_compress_sin,
|
|
metadata.query_start_loc,
|
|
metadata.start_pos,
|
|
metadata.block_table,
|
|
metadata.block_size,
|
|
DeviceOperator.get_dsa_compressor_slot_mapping_format(),
|
|
self.compress_ratio,
|
|
metadata.num_compressed_tokens,
|
|
metadata.num_reqs_actual,
|
|
)
|
|
|
|
def process_weights_after_loading(self, act_dtype: torch.dtype):
|
|
if self.attn_sink.numel() != self.num_heads:
|
|
raise RuntimeError(
|
|
"DSA-CP expects full-head attn_sink loaded on every TP rank, "
|
|
f"got {self.attn_sink.numel()} heads, expected {self.num_heads}."
|
|
)
|
|
if self.enable_dsa_cp_with_o_proj_tp:
|
|
self._maybe_init_o_proj_tp_full_params()
|
|
|
|
@staticmethod
|
|
def _check_dynamic_quant(layer: torch.nn.Module) -> bool:
|
|
return get_ascend_device_type() in {AscendDeviceType.A5} and hasattr(layer, "weight_scale")
|
|
|
|
def _maybe_init_o_proj_tp_full_params(self) -> None:
|
|
self._wo_a_dynamic_quant = type(self)._check_dynamic_quant(self.wo_a)
|
|
self._wo_b_dynamic_quant = type(self)._check_dynamic_quant(self.wo_b)
|
|
if AscendDSACPImpl.wo_a_full_pool is None:
|
|
sample = self.wo_a.weight
|
|
AscendDSACPImpl.wo_a_full_pool = torch.empty(
|
|
(sample.shape[0] * self.tp_size, *sample.shape[1:]),
|
|
dtype=sample.dtype,
|
|
device=sample.device,
|
|
)
|
|
self.wo_a_tp_weight = self.wo_a.weight.clone().detach().contiguous()
|
|
self.wo_a.weight.set_(self.wo_a_tp_weight)
|
|
if AscendDSACPImpl.wo_b_full_pool is None:
|
|
sample = self.wo_b.weight
|
|
AscendDSACPImpl.wo_b_full_pool = torch.empty(
|
|
(sample.shape[0] * self.tp_size, *sample.shape[1:]),
|
|
dtype=sample.dtype,
|
|
device=sample.device,
|
|
)
|
|
self.wo_b_tp_weight = self.wo_b.weight.clone().detach().contiguous()
|
|
self.wo_b.weight.set_(self.wo_b_tp_weight)
|
|
|
|
if self._wo_a_dynamic_quant:
|
|
if AscendDSACPImpl.wo_a_full_weight_scale_pool is None:
|
|
sample = self.wo_a.weight_scale
|
|
AscendDSACPImpl.wo_a_full_weight_scale_pool = torch.empty(
|
|
(sample.shape[0] * self.tp_size, *sample.shape[1:]),
|
|
dtype=sample.dtype,
|
|
device=sample.device,
|
|
)
|
|
self.wo_a_tp_weight_scale = self.wo_a.weight_scale.clone().detach().contiguous()
|
|
self.wo_a.weight_scale.set_(self.wo_a_tp_weight_scale)
|
|
if self._wo_b_dynamic_quant:
|
|
if AscendDSACPImpl.wo_b_full_weight_scale_pool is None:
|
|
sample = self.wo_b.weight_scale
|
|
AscendDSACPImpl.wo_b_full_weight_scale_pool = torch.empty(
|
|
(sample.shape[0] * self.tp_size, *sample.shape[1:]),
|
|
dtype=sample.dtype,
|
|
device=sample.device,
|
|
)
|
|
self.wo_b_tp_weight_scale = self.wo_b.weight_scale.clone().detach().contiguous()
|
|
self.wo_b.weight_scale.set_(self.wo_b_tp_weight_scale)
|
|
|
|
def _maybe_all_gather_o_proj_full_weight(
|
|
self,
|
|
enabled: bool,
|
|
) -> list[torch.distributed.Work]:
|
|
if not enabled:
|
|
return []
|
|
handles = []
|
|
assert AscendDSACPImpl.wo_a_full_pool is not None
|
|
_, weight_handle = all_gather_async(
|
|
self.wo_a_tp_weight,
|
|
self.tp_group,
|
|
output=AscendDSACPImpl.wo_a_full_pool,
|
|
)
|
|
if weight_handle is not None:
|
|
handles.append(weight_handle)
|
|
assert AscendDSACPImpl.wo_b_full_pool is not None
|
|
_, wo_b_weight_handle = all_gather_async(
|
|
self.wo_b_tp_weight,
|
|
self.tp_group,
|
|
output=AscendDSACPImpl.wo_b_full_pool,
|
|
)
|
|
if wo_b_weight_handle is not None:
|
|
handles.append(wo_b_weight_handle)
|
|
if self._wo_a_dynamic_quant:
|
|
assert AscendDSACPImpl.wo_a_full_weight_scale_pool is not None
|
|
_, weight_scale_handle = all_gather_async(
|
|
self.wo_a_tp_weight_scale,
|
|
self.tp_group,
|
|
output=AscendDSACPImpl.wo_a_full_weight_scale_pool,
|
|
)
|
|
if weight_scale_handle is not None:
|
|
handles.append(weight_scale_handle)
|
|
if self._wo_b_dynamic_quant:
|
|
assert AscendDSACPImpl.wo_b_full_weight_scale_pool is not None
|
|
_, wo_b_weight_scale_handle = all_gather_async(
|
|
self.wo_b_tp_weight_scale,
|
|
self.tp_group,
|
|
output=AscendDSACPImpl.wo_b_full_weight_scale_pool,
|
|
)
|
|
if wo_b_weight_scale_handle is not None:
|
|
handles.append(wo_b_weight_scale_handle)
|
|
return handles
|
|
|
|
def _switch_o_proj_to_full_weight(
|
|
self,
|
|
handles: list[torch.distributed.Work],
|
|
) -> None:
|
|
for handle in handles:
|
|
handle.wait()
|
|
assert AscendDSACPImpl.wo_a_full_pool is not None
|
|
self.wo_a.weight.set_(AscendDSACPImpl.wo_a_full_pool)
|
|
if self._wo_a_dynamic_quant:
|
|
assert AscendDSACPImpl.wo_a_full_weight_scale_pool is not None
|
|
self.wo_a.weight_scale.set_(AscendDSACPImpl.wo_a_full_weight_scale_pool)
|
|
assert AscendDSACPImpl.wo_b_full_pool is not None
|
|
self.wo_b.weight.set_(AscendDSACPImpl.wo_b_full_pool)
|
|
if self._wo_b_dynamic_quant:
|
|
assert AscendDSACPImpl.wo_b_full_weight_scale_pool is not None
|
|
self.wo_b.weight_scale.set_(AscendDSACPImpl.wo_b_full_weight_scale_pool)
|
|
|
|
def _switch_o_proj_to_tp_weight(self) -> None:
|
|
self.wo_a.weight.set_(self.wo_a_tp_weight)
|
|
if self._wo_a_dynamic_quant:
|
|
self.wo_a.weight_scale.set_(self.wo_a_tp_weight_scale)
|
|
self.wo_b.weight.set_(self.wo_b_tp_weight)
|
|
if self._wo_b_dynamic_quant:
|
|
self.wo_b.weight_scale.set_(self.wo_b_tp_weight_scale)
|
|
|
|
def _apply_wo_b(
|
|
self,
|
|
o_proj_input: torch.Tensor,
|
|
full_weight: bool,
|
|
) -> torch.Tensor:
|
|
if not full_weight:
|
|
return self.wo_b(o_proj_input)
|
|
return self.wo_b.quant_method.apply(self.wo_b, o_proj_input, bias=None)
|
|
|
|
def forward( # type: ignore[override]
|
|
self,
|
|
layer_name,
|
|
hidden_states: torch.Tensor, # query in unified attn
|
|
kv_cache: tuple[torch.Tensor],
|
|
attn_metadata: list[M],
|
|
need_gather_q_kv: bool = False,
|
|
output: torch.Tensor | None = None,
|
|
) -> torch.Tensor:
|
|
assert output is not None, "Output tensor must be provided."
|
|
if attn_metadata is None:
|
|
# Profiling run.
|
|
return output.fill_(0)
|
|
if not isinstance(attn_metadata, list):
|
|
attn_metadata = [attn_metadata]
|
|
wait_for_kv_layer_from_connector(layer_name)
|
|
full_gather_wo_a_enabled = (
|
|
self.tp_size > 1
|
|
and self.enable_dsa_cp_with_o_proj_tp
|
|
and attn_metadata[0].attn_state
|
|
not in {
|
|
AscendAttentionState.DecodeOnly,
|
|
AscendAttentionState.SpecDecoding,
|
|
}
|
|
)
|
|
local_attn_output, o_proj_full_handles = self._forward(
|
|
layer_name,
|
|
hidden_states,
|
|
kv_cache,
|
|
attn_metadata,
|
|
need_gather_q_kv,
|
|
full_gather_wo_a_enabled,
|
|
)
|
|
o_proj_input = self._restore_tp_head_layout(
|
|
local_attn_output,
|
|
layer_name,
|
|
attn_metadata[0],
|
|
skip_all_to_all=full_gather_wo_a_enabled,
|
|
)
|
|
num_tokens = o_proj_input.shape[0]
|
|
|
|
# o
|
|
if full_gather_wo_a_enabled:
|
|
self._switch_o_proj_to_full_weight(o_proj_full_handles)
|
|
o_proj_groups = self.n_group if full_gather_wo_a_enabled else self.n_local_groups
|
|
try:
|
|
if get_ascend_device_type() in {AscendDeviceType.A5}:
|
|
o = o_proj_input.view(num_tokens, o_proj_groups, -1)
|
|
o, swiglu_out_scale = torch_npu.npu_dynamic_mx_quant(o, dst_type=torch.float8_e4m3fn)
|
|
o = torch_npu.npu_transpose_quant_batchmatmul(
|
|
o,
|
|
self.wo_a.weight,
|
|
dtype=torch.bfloat16,
|
|
bias=None,
|
|
group_sizes=(0, 0, 32),
|
|
x1_scale=swiglu_out_scale.view(torch.float8_e8m0fnu),
|
|
x2_scale=self.wo_a.weight_scale.view(torch.float8_e8m0fnu),
|
|
perm_x1=(1, 0, 2),
|
|
perm_x2=(0, 1, 2),
|
|
perm_y=(1, 0, 2),
|
|
)
|
|
o = o.reshape(num_tokens, -1)
|
|
output[...] = self._apply_wo_b(o, full_gather_wo_a_enabled)
|
|
else:
|
|
o_proj_input = o_proj_input.view(num_tokens, o_proj_groups, -1)
|
|
if olora_tp_enable():
|
|
o_proj_input = self.wo_a(o_proj_input)
|
|
else:
|
|
# wo_a = self.wo_a.weight.view(o_proj_groups, self.o_lora_rank, -1)
|
|
# o = torch.einsum("tgd,grd->tgr", o, wo_a)
|
|
o_proj_input = torch_npu.npu_transpose_batchmatmul(
|
|
o_proj_input,
|
|
self.wo_a.weight,
|
|
bias=None,
|
|
scale=None,
|
|
perm_x1=(1, 0, 2),
|
|
perm_x2=(0, 1, 2),
|
|
perm_y=(1, 0, 2),
|
|
batch_split_factor=1,
|
|
)
|
|
o_proj_input = o_proj_input.reshape(num_tokens, -1)
|
|
output[...] = self._apply_wo_b(o_proj_input, full_gather_wo_a_enabled)
|
|
finally:
|
|
if full_gather_wo_a_enabled:
|
|
self._switch_o_proj_to_tp_weight()
|
|
|
|
maybe_save_kv_layer_to_connector(layer_name, list(kv_cache))
|
|
|
|
return output
|
|
|
|
def _forward(
|
|
self,
|
|
layer_name,
|
|
hidden_states_local: torch.Tensor,
|
|
kv_cache: tuple,
|
|
attn_metadata: list[M],
|
|
need_gather_q_kv: bool = False,
|
|
full_gather_wo_a_enabled: bool = False,
|
|
):
|
|
"""Run full-sequence KV cache updates and local-token attention."""
|
|
(compress_kv_cache, swa_kv_cache, state_cache, _, _, _) = DeviceOperator.unpack_dsa_forward_kv_cache(
|
|
kv_cache, self.compress_ratio
|
|
)
|
|
if self.compress_ratio == 4:
|
|
(compressor_attn_metadata, compressor_kv_state_metadata, _, _, swa_metadata) = attn_metadata
|
|
elif self.compress_ratio == 128:
|
|
(compressor_attn_metadata, compressor_kv_state_metadata, swa_metadata) = attn_metadata
|
|
else:
|
|
(swa_metadata,) = attn_metadata
|
|
common_attn_metadata = attn_metadata[0]
|
|
|
|
hidden_states = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(hidden_states_local, need_gather_q_kv)
|
|
|
|
assert common_attn_metadata.req_metadata is not None
|
|
assert swa_metadata.req_metadata is not None
|
|
req_metadata = common_attn_metadata.req_metadata
|
|
cp_metadata = req_metadata.cp_metadata
|
|
cos = req_metadata.cos[layer_name]
|
|
sin = req_metadata.sin[layer_name]
|
|
local_cos = cp_metadata.local_cos[layer_name]
|
|
local_sin = cp_metadata.local_sin[layer_name]
|
|
actual_seq_lengths_query = req_metadata.query_start_loc
|
|
local_seq_lengths_query = cp_metadata.local_query_start_loc
|
|
local_seq_lengths_key = cp_metadata.local_seq_lens
|
|
has_prefill = _has_prefill(common_attn_metadata.attn_state)
|
|
hidden_states_cache = hidden_states[: common_attn_metadata.num_actual_tokens]
|
|
|
|
if (not isinstance(self.wq_b.quant_method, AscendUnquantizedLinearMethod)) and isinstance(
|
|
self.wq_b.quant_method.quant_method, AscendW8A8DynamicLinearMethod
|
|
):
|
|
q_a = self.wq_a(hidden_states_local)
|
|
qr_local, qr_pertoken_scale_local = torch.ops._C_ascend.npu_rms_norm_dynamic_quant(
|
|
q_a, self.q_norm.weight, epsilon=self.eps
|
|
)
|
|
if getattr(self.wq_b, "_chunk_size", 0):
|
|
bias = self.wq_b.bias
|
|
chunk_size = self.wq_b._chunk_size
|
|
bias_1 = bias[:chunk_size] if bias is not None else None
|
|
bias_2 = bias[chunk_size:] if bias is not None else None
|
|
q = torch.cat(
|
|
[
|
|
torch_npu.npu_quant_matmul(
|
|
qr_local,
|
|
self.wq_b.weight_1,
|
|
self.wq_b.weight_1_scale,
|
|
pertoken_scale=qr_pertoken_scale_local,
|
|
bias=bias_1,
|
|
output_dtype=hidden_states_local.dtype,
|
|
),
|
|
torch_npu.npu_quant_matmul(
|
|
qr_local,
|
|
self.wq_b.weight_2,
|
|
self.wq_b.weight_2_scale,
|
|
pertoken_scale=qr_pertoken_scale_local,
|
|
bias=bias_2,
|
|
output_dtype=hidden_states_local.dtype,
|
|
),
|
|
],
|
|
dim=-1,
|
|
)
|
|
else:
|
|
q = torch_npu.npu_quant_matmul(
|
|
qr_local,
|
|
self.wq_b.weight,
|
|
self.wq_b.weight_scale,
|
|
pertoken_scale=qr_pertoken_scale_local,
|
|
bias=self.wq_b.bias,
|
|
output_dtype=hidden_states_local.dtype,
|
|
)
|
|
else:
|
|
qr_local = self.q_norm(self.wq_a(hidden_states_local))
|
|
q = self.wq_b(qr_local)
|
|
qr_pertoken_scale_local = None
|
|
|
|
q = q.unflatten(-1, (self.num_heads, self.head_dim))
|
|
|
|
q = DeviceOperator.apply_dsa_q_rms(q, self.eps, self.q_norm_without_weight)
|
|
torch.ops._C_ascend.inplace_partial_rotary_mul(
|
|
q.unsqueeze(1),
|
|
local_cos,
|
|
local_sin,
|
|
rotary_mode="interleave",
|
|
partial_slice=[self.nope_head_dim, self.head_dim],
|
|
)
|
|
|
|
o_proj_full_handles = self._maybe_all_gather_o_proj_full_weight(full_gather_wo_a_enabled)
|
|
|
|
kv = self.wkv(hidden_states_cache)
|
|
kv = self.kv_norm(kv)
|
|
assert self.rope_head_dim is not None
|
|
kv = kv.view(-1, 1, self.nope_head_dim + self.rope_head_dim)
|
|
torch.ops._C_ascend.inplace_partial_rotary_mul(
|
|
kv.unsqueeze(1),
|
|
cos[: kv.shape[0]],
|
|
sin[: kv.shape[0]],
|
|
rotary_mode="interleave",
|
|
partial_slice=[self.nope_head_dim, self.head_dim],
|
|
)
|
|
DeviceOperator.dsa_kv_compress_scatter(swa_kv_cache, kv, swa_metadata.req_metadata.slot_mapping)
|
|
|
|
compress_topk_idxs = None
|
|
if self.compress_ratio > 1:
|
|
assert compressor_attn_metadata.req_metadata is not None
|
|
assert compressor_kv_state_metadata.req_metadata is not None
|
|
if self.compress_ratio == 4:
|
|
self._update_indexer_cache(
|
|
x=hidden_states_cache,
|
|
kv_cache=kv_cache,
|
|
attn_metadata=attn_metadata,
|
|
actual_seq_lengths_query=actual_seq_lengths_query,
|
|
)
|
|
compress_topk_idxs = self._indexer_select_topk(
|
|
x=hidden_states_local,
|
|
qr=qr_local,
|
|
kv_cache=kv_cache,
|
|
attn_metadata=attn_metadata,
|
|
cos=local_cos,
|
|
sin=local_sin,
|
|
actual_seq_lengths_query=local_seq_lengths_query,
|
|
actual_seq_lengths_key=local_seq_lengths_key,
|
|
qr_pertoken_scale=qr_pertoken_scale_local,
|
|
)
|
|
|
|
coff = 2 if self.compressor_overlap else 1
|
|
compress_cos, compress_sin, compress_slot_mapping = self._compute_compressor_metadata(
|
|
compressor_attn_metadata.req_metadata,
|
|
)
|
|
compressed_kv = torch.ops._C_ascend.compressor(
|
|
hidden_states_cache,
|
|
self.compressor_wkv.weight,
|
|
self.compressor_wgate.weight,
|
|
state_cache.squeeze(-2),
|
|
self.compressor_ape,
|
|
self.compressor_norm.weight,
|
|
compress_sin.view(-1, compress_sin.shape[-1]),
|
|
compress_cos.view(-1, compress_cos.shape[-1]),
|
|
state_block_table=compressor_kv_state_metadata.req_metadata.block_table,
|
|
cu_seqlens=actual_seq_lengths_query,
|
|
seqused=None,
|
|
start_pos=req_metadata.start_pos,
|
|
rope_head_dim=self.rope_head_dim,
|
|
cmp_ratio=self.compress_ratio,
|
|
coff=coff,
|
|
norm_eps=self.compressor_norm_eps,
|
|
rotary_mode=2,
|
|
cache_mode=1,
|
|
)
|
|
|
|
if compressed_kv.numel() == 0:
|
|
compressed_kv = None
|
|
DeviceOperator.dsa_kv_compress_scatter(compress_kv_cache, compressed_kv, compress_slot_mapping)
|
|
|
|
notify_kv_cache_written(layer_name)
|
|
record_attention_compute_start()
|
|
attn_op = DeviceOperator.get_dsa_sparse_attn_op()
|
|
extra_attn_kwargs: dict = DeviceOperator.get_dsa_sparse_attn_base_kwargs()
|
|
if has_prefill:
|
|
DeviceOperator.add_dsa_sparse_attn_extra_kwargs(
|
|
extra_attn_kwargs, cu_seqlens_ori_kv=local_seq_lengths_query
|
|
)
|
|
|
|
common_attn_kwargs = dict(
|
|
cu_seqlens_q=local_seq_lengths_query,
|
|
seqused_kv=local_seq_lengths_key,
|
|
sinks=self.attn_sink,
|
|
softmax_scale=self.softmax_scale,
|
|
cmp_ratio=max(self.compress_ratio, 1),
|
|
ori_mask_mode=4,
|
|
ori_win_left=self.window_size - 1,
|
|
ori_win_right=0,
|
|
layout_q="TND",
|
|
layout_kv="PA_ND",
|
|
**extra_attn_kwargs,
|
|
)
|
|
|
|
if self.compress_ratio <= 1:
|
|
attn_output = attn_op(
|
|
q,
|
|
ori_kv=swa_kv_cache,
|
|
ori_block_table=swa_metadata.req_metadata.block_table,
|
|
metadata=swa_metadata.req_metadata.sas_metadata,
|
|
**common_attn_kwargs,
|
|
)[0]
|
|
elif self.compress_ratio == 4:
|
|
assert compressor_attn_metadata.req_metadata is not None
|
|
DeviceOperator.add_dsa_sparse_attn_extra_kwargs(
|
|
common_attn_kwargs, cu_seqlens_cmp_kv=req_metadata.cu_cmp_seqlen_list
|
|
)
|
|
attn_output = attn_op(
|
|
q,
|
|
ori_kv=swa_kv_cache,
|
|
cmp_kv=compress_kv_cache,
|
|
cmp_sparse_indices=compress_topk_idxs,
|
|
ori_block_table=swa_metadata.req_metadata.block_table,
|
|
cmp_block_table=compressor_attn_metadata.req_metadata.block_table,
|
|
metadata=req_metadata.sas_metadata,
|
|
cmp_mask_mode=3,
|
|
**common_attn_kwargs,
|
|
)[0]
|
|
else:
|
|
assert compressor_attn_metadata.req_metadata is not None
|
|
DeviceOperator.add_dsa_sparse_attn_extra_kwargs(
|
|
common_attn_kwargs, cu_seqlens_cmp_kv=req_metadata.cu_cmp_seqlen_list
|
|
)
|
|
attn_output = attn_op(
|
|
q,
|
|
ori_kv=swa_kv_cache,
|
|
cmp_kv=compress_kv_cache,
|
|
ori_block_table=swa_metadata.req_metadata.block_table,
|
|
cmp_block_table=compressor_attn_metadata.req_metadata.block_table,
|
|
metadata=compressor_attn_metadata.req_metadata.sas_metadata,
|
|
cmp_mask_mode=3,
|
|
**common_attn_kwargs,
|
|
)[0]
|
|
return attn_output, o_proj_full_handles
|
|
|
|
def _restore_tp_head_layout(
|
|
self,
|
|
local_attn_output: torch.Tensor,
|
|
layer_name: str,
|
|
attn_metadata: M,
|
|
skip_all_to_all: bool = False,
|
|
) -> torch.Tensor:
|
|
assert attn_metadata.req_metadata is not None
|
|
req_metadata = attn_metadata.req_metadata
|
|
cp_metadata = req_metadata.cp_metadata
|
|
num_tokens = local_attn_output.shape[0]
|
|
torch.ops._C_ascend.inplace_partial_rotary_mul(
|
|
local_attn_output.unsqueeze(1),
|
|
cp_metadata.local_cos[layer_name],
|
|
-cp_metadata.local_sin[layer_name],
|
|
rotary_mode="interleave",
|
|
partial_slice=[self.nope_head_dim, self.head_dim],
|
|
)
|
|
|
|
if self.tp_size == 1 or skip_all_to_all:
|
|
return local_attn_output
|
|
|
|
send = (
|
|
local_attn_output.view(num_tokens, self.tp_size, self.n_local_heads, self.head_dim)
|
|
.permute(1, 0, 2, 3)
|
|
.contiguous()
|
|
.view(-1, self.n_local_heads, self.head_dim)
|
|
)
|
|
recv = torch.empty_like(send)
|
|
dist.all_to_all_single(recv, send, group=self.tp_group.device_group)
|
|
return recv
|
|
|
|
def _update_indexer_cache(
|
|
self,
|
|
x: torch.Tensor,
|
|
kv_cache: tuple[torch.Tensor, ...],
|
|
attn_metadata: list[M],
|
|
actual_seq_lengths_query: torch.Tensor,
|
|
) -> None:
|
|
(indexer_state_cache, indexer_k_cache, indexer_scale_cache, indexer_full_cache) = (
|
|
DeviceOperator.unpack_dsa_indexer_kv_cache(kv_cache)
|
|
)
|
|
(_, _, indexer_kv_state_metadata, indexer_kv_scale_metadata, _) = attn_metadata
|
|
coff = 2 if self.compressor_overlap else 1
|
|
assert indexer_kv_scale_metadata is not None
|
|
assert indexer_kv_state_metadata is not None
|
|
assert indexer_kv_scale_metadata.req_metadata is not None
|
|
assert indexer_kv_state_metadata.req_metadata is not None
|
|
assert self.indexer is not None
|
|
compressed_cos, compressed_sin, indexer_slot_mapping = self._compute_compressor_metadata(
|
|
indexer_kv_scale_metadata.req_metadata,
|
|
)
|
|
kv = torch.ops._C_ascend.compressor(
|
|
x,
|
|
self.indexcom_wkv.weight,
|
|
self.indexcom_wgate.weight,
|
|
indexer_state_cache.squeeze(-2),
|
|
self.indexcom_ape,
|
|
self.indexcom_norm.weight,
|
|
compressed_sin.view(-1, compressed_sin.shape[-1]),
|
|
compressed_cos.view(-1, compressed_cos.shape[-1]),
|
|
state_block_table=indexer_kv_state_metadata.req_metadata.block_table,
|
|
cu_seqlens=actual_seq_lengths_query,
|
|
seqused=None,
|
|
start_pos=indexer_kv_scale_metadata.req_metadata.start_pos,
|
|
rope_head_dim=self.rope_head_dim,
|
|
cmp_ratio=self.compress_ratio,
|
|
coff=coff,
|
|
norm_eps=self.compressor_norm_eps,
|
|
rotary_mode=2,
|
|
cache_mode=1,
|
|
)
|
|
|
|
if kv.numel() == 0:
|
|
return
|
|
if self.indexer.compressor.rotate:
|
|
kv = rotate_activation(kv, indexer_kv_scale_metadata.hadamard)
|
|
|
|
_, kv_scale = DeviceOperator.indexer_quant_scatter_part1(
|
|
kv,
|
|
indexer_k_cache,
|
|
indexer_full_cache,
|
|
indexer_slot_mapping,
|
|
)
|
|
if kv_scale is not None:
|
|
DeviceOperator.dsa_indexer_scatter_scale_part3(
|
|
kv_scale,
|
|
indexer_scale_cache,
|
|
indexer_slot_mapping,
|
|
)
|
|
|
|
def _indexer_select_topk(
|
|
self,
|
|
x: torch.Tensor,
|
|
qr: torch.Tensor,
|
|
kv_cache: tuple[torch.Tensor, ...],
|
|
attn_metadata: list[M],
|
|
cos: torch.Tensor,
|
|
sin: torch.Tensor,
|
|
actual_seq_lengths_query: torch.Tensor,
|
|
actual_seq_lengths_key: torch.Tensor,
|
|
qr_pertoken_scale: torch.Tensor = None,
|
|
):
|
|
(_, indexer_k_cache, indexer_scale_cache, _) = DeviceOperator.unpack_dsa_indexer_kv_cache(kv_cache)
|
|
(_, _, _, indexer_kv_scale_metadata, _) = attn_metadata
|
|
assert indexer_kv_scale_metadata is not None
|
|
|
|
if (
|
|
(not isinstance(self.inderxer_wq_b.quant_method, AscendUnquantizedLinearMethod))
|
|
and isinstance(self.inderxer_wq_b.quant_method.quant_method, AscendW8A8DynamicLinearMethod)
|
|
and qr_pertoken_scale is not None
|
|
and get_ascend_device_type() not in {AscendDeviceType.A5}
|
|
):
|
|
q = torch_npu.npu_quant_matmul(
|
|
qr,
|
|
self.inderxer_wq_b.weight,
|
|
self.inderxer_wq_b.weight_scale,
|
|
pertoken_scale=qr_pertoken_scale,
|
|
bias=self.inderxer_wq_b.bias,
|
|
output_dtype=x.dtype,
|
|
)
|
|
else:
|
|
q = self.inderxer_wq_b(qr)
|
|
q = q.view(-1, self.indexer_heads, self.indexcom_head_dim)
|
|
torch.ops._C_ascend.inplace_partial_rotary_mul(
|
|
q.unsqueeze(1),
|
|
cos,
|
|
sin,
|
|
rotary_mode="interleave",
|
|
partial_slice=[self.indexcom_head_dim - self.rope_head_dim, self.indexcom_head_dim],
|
|
)
|
|
q = rotate_activation(q, indexer_kv_scale_metadata.hadamard)
|
|
weights = self.weights_proj(x) * (self.indexer_softmax_scale * self.indexer_heads**-0.5)
|
|
|
|
q, q_scale = DeviceOperator.indexer_quantize_query(q)
|
|
|
|
assert indexer_kv_scale_metadata.req_metadata is not None
|
|
qli_metadata = indexer_kv_scale_metadata.req_metadata.qli_metadata
|
|
block_table = indexer_kv_scale_metadata.req_metadata.block_table
|
|
topk_idxs, _ = torch.ops._C_ascend.npu_vllm_quant_lightning_indexer(
|
|
query=q,
|
|
key=indexer_k_cache,
|
|
weights=DeviceOperator.prepare_dsa_indexer_weights(weights),
|
|
query_dequant_scale=DeviceOperator.prepare_dsa_indexer_query_scale(q_scale),
|
|
key_dequant_scale=DeviceOperator.prepare_dsa_indexer_key_scale(indexer_scale_cache),
|
|
actual_seq_lengths_query=actual_seq_lengths_query[1:],
|
|
actual_seq_lengths_key=actual_seq_lengths_key,
|
|
block_table=block_table,
|
|
metadata=qli_metadata,
|
|
query_quant_mode=0,
|
|
key_quant_mode=0,
|
|
layout_query="TND",
|
|
layout_key="PA_BSND",
|
|
sparse_count=self.index_topk,
|
|
sparse_mode=3,
|
|
pre_tokens=(1 << 63) - 1,
|
|
next_tokens=(1 << 63) - 1,
|
|
cmp_ratio=4,
|
|
return_value=False,
|
|
)
|
|
return topk_idxs
|
|
|
|
def dsa_warmup_with_multistream(self, hidden_states: torch.Tensor):
|
|
pass
|