diff --git a/dlinfer/framework/lmdeploy_ext/cudagraph/ascend_cudagraph.py b/dlinfer/framework/lmdeploy_ext/cudagraph/ascend_cudagraph.py index a16d2a50..baf75afa 100644 --- a/dlinfer/framework/lmdeploy_ext/cudagraph/ascend_cudagraph.py +++ b/dlinfer/framework/lmdeploy_ext/cudagraph/ascend_cudagraph.py @@ -1,41 +1,24 @@ # Copyright (c) 2024, OpenMMLab and DeepLink. All rights reserved. # this file implements the cudagraph for ascend backend. import functools -from typing import Any, Dict, List, Optional -from dataclasses import dataclass from contextlib import ExitStack -from packaging.version import InvalidVersion, Version +from typing import Any, Dict, List import torch -import torch_npu from torch import Tensor from torch.profiler import record_function -from lmdeploy.pytorch.model_inputs import get_step_ctx_manager -from lmdeploy.pytorch.models.utils.cudagraph import CudaGraphMeta -from lmdeploy.pytorch.models.utils.cudagraph import CudaGraphMixin +from lmdeploy.pytorch.backends.graph_runner import GraphRunner from lmdeploy.pytorch.config import BackendConfig, CacheConfig, ModelConfig from lmdeploy.pytorch.model_inputs import StepContext, get_step_ctx_manager -from lmdeploy.pytorch.backends.graph_runner import GraphRunner - +from lmdeploy.pytorch.models.utils.cudagraph import CudaGraphMeta +from lmdeploy.pytorch.models.utils.cudagraph import CudaGraphMixin from lmdeploy.utils import get_logger logger = get_logger("dlinfer") BuffType = Dict[str, Tensor] -@functools.lru_cache() -def aclgraph_use_torch_npu_update(): - min_valid_version = Version("2.8.0.post1") - - try: - current_version = Version(torch_npu.__version__) - except InvalidVersion: - return False - - return current_version >= min_valid_version - - # AscendCudaGraphMixin methods for cudagraph buffer management. def AscendCudaGraphMixin_support_cuda_graph( self, @@ -46,7 +29,7 @@ def AscendCudaGraphMixin_support_cuda_graph( inputs_embeds: Tensor = None, **kwargs, ): - """Allow multi-token decode graph only when runtime length updates exist.""" + """Allow decode graph after the Ascend runtime was validated at import.""" if attn_metadata is None: return False @@ -55,11 +38,9 @@ def AscendCudaGraphMixin_support_cuda_graph( # (any rank prefill => all prefill) rather than this rank's local is_decoding. if not get_step_ctx_manager().current_context().global_is_decoding(): return False - is_decoding = getattr(attn_metadata, "is_decoding", False) - is_multi_token = getattr(attn_metadata, "is_multi_token_decoding", False) - if is_multi_token and not aclgraph_use_torch_npu_update(): - return False - return is_decoding or is_multi_token + return getattr(attn_metadata, "is_decoding", False) or getattr( + attn_metadata, "is_multi_token_decoding", False + ) def AscendCudaGraphMixin_make_buffers_cudagraph( @@ -309,7 +290,6 @@ def _get_capture_batch_size_impl(max_batches: int): if max_batches not in ret: ret.append(max_batches) - set_graph_params(set(ret)) return ret @@ -331,7 +311,7 @@ def __init__( pool: Any, model_config: ModelConfig, device: torch.device, - update_stream: torch.npu.Stream, + is_mla: bool, ): self.model = model self.ctx_mgr = model.ctx_mgr @@ -356,7 +336,7 @@ def __init__( self.is_decoding = is_decoding self.pool = pool self._graph: torch.npu.NPUGraph = None - self.update_stream = update_stream + self.is_mla = is_mla @record_function("capture_cudagraph") def capture(self, **kwargs): @@ -373,15 +353,17 @@ def capture(self, **kwargs): warmup_buffers = self.model.make_output_buffers(warmup_output) aclgraph = torch.npu.NPUGraph() - with ExitStack() as stack: - AscendGraphRunner.capturing = True - with torch.npu.graph( - aclgraph, - auto_dispatch_capture=True, - pool=self.pool, - stream=current_stream, - ): - graph_output = self.model(**padded_kwargs) + AscendGraphRunner.capturing = True + try: + with ExitStack(): + with torch.npu.graph( + aclgraph, + auto_dispatch_capture=True, + pool=self.pool, + stream=current_stream, + ): + graph_output = self.model(**padded_kwargs) + finally: AscendGraphRunner.capturing = False output_buffers = self.model.make_output_buffers(graph_output) @@ -397,16 +379,16 @@ def forward(self, **kwargs): self.model.fill_buffers_cudagraph(self.meta, **kwargs) context = self.ctx_mgr.current_context() self.model.update_context_cudagraph(self.meta, context) - if aclgraph_use_torch_npu_update(): - self._graph.replay() - self._graph.update( - cpu_update_input=[ - {"actual_seq_lengths_kv": self.meta.input_buffers["kv_seqlens"]} - ] - ) + self._graph.replay() + if self.is_mla: + cpu_update_input = [ + {"actual_seq_kvlen": self.meta.input_buffers["kv_seqlens"].tolist()} + ] else: - update_attn_params(self.update_stream, self.meta, self.max_batches) - self._graph.replay() + cpu_update_input = [ + {"actual_seq_lengths_kv": self.meta.input_buffers["kv_seqlens"]} + ] + self._graph.update(cpu_update_input=cpu_update_input) output_buffers = self.meta.output_buffers output = self.model.get_outputs_cudagraph(output_buffers, **kwargs) return output @@ -445,6 +427,7 @@ def __init__( cache_config: CacheConfig, backend_config: BackendConfig, device: torch.device, + is_mla: bool = False, ): super().__init__(model, model_config, cache_config, backend_config, device) self.max_batches = cache_config.max_batches @@ -454,7 +437,7 @@ def __init__( self.graph_pool_handle = torch.cuda.graph_pool_handle() self._runner_map: Dict[Any, AscendSingleGraphRunner] = dict() self.has_try_compile_model: bool = False - self.update_stream = torch.npu.Stream() + self.is_mla = is_mla def check_enable_graph(self): """Check enable graph.""" @@ -535,7 +518,7 @@ def __call__(self, **kwargs): pool=self.graph_pool_handle, model_config=self.model_config, device=self.device, - update_stream=self.update_stream, + is_mla=self.is_mla, ) runner.capture(**kwargs) self._runner_map[graph_key] = runner @@ -560,13 +543,13 @@ def prepare_inputs_for_generation( def reset(self): """Remove all graphs and related resources to prevent hanging on exit.""" + super().reset() for _, runner in self._runner_map.items(): try: runner.reset() except Exception as e: logger.warning(f"AscendGraphRunner.reset: runner.reset error: {e!r}") self._runner_map.clear() - clear_graph_params() self.graph_pool_handle = None torch.npu.empty_cache() @@ -595,140 +578,6 @@ def update_inputs(self, inputs): def get_capture_batch_sizes(self) -> List[int]: """Capture batch sizes.""" + if self.cache_config.cudagraph_capture_batch_sizes is not None: + return super().get_capture_batch_sizes() return _get_capture_batch_size_impl(self.cache_config.max_batches) - - -@dataclass -class GraphParams: - events: dict[int, list[torch.npu.ExternalEvent]] - workspaces: dict[int, torch.Tensor] - handles: dict[int, list[torch_npu._C._NPUTaskGroupHandle]] - attn_params: dict[int, list[tuple]] - is_mla: bool - - -_graph_params: Optional[GraphParams] = None -_graph_capture_sizes: set[int] = None - - -def set_graph_params(aclgraph_capture_sizes: set[int]): - global _graph_params - global _graph_capture_sizes - if _graph_params is not None: - raise ValueError("Graph parameters have already been set!") - _graph_params = GraphParams( - events={size: [] for size in aclgraph_capture_sizes}, - workspaces={size: None for size in aclgraph_capture_sizes}, - handles={size: [] for size in aclgraph_capture_sizes}, - attn_params={size: [] for size in aclgraph_capture_sizes}, - is_mla=False, - ) - _graph_capture_sizes = aclgraph_capture_sizes - - -def get_graph_params(): - return _graph_params - - -def clear_graph_params(): - """Clear global graph params and release references to KV cache tensors.""" - global _graph_params - global _graph_capture_sizes - if _graph_params is None: - return - - try: - for k in list(_graph_params.attn_params.keys()): - _graph_params.attn_params[k].clear() - for k in list(_graph_params.handles.keys()): - _graph_params.handles[k].clear() - for k in list(_graph_params.events.keys()): - _graph_params.events[k].clear() - _graph_params.is_mla = None - - _graph_params.workspaces.clear() - finally: - _graph_params = None - _graph_capture_sizes = None - # 清除 lru_cache,使下次推理时 _get_capture_batch_size_impl - # 重新执行并调用 set_graph_params 干净重建 - _get_capture_batch_size_impl.cache_clear() - - -def update_attn_params(update_stream, forward_meta, runtime_size): - graph_params = get_graph_params() - for param, handle, event in zip( - graph_params.attn_params[runtime_size], - graph_params.handles[runtime_size], - graph_params.events[runtime_size], - ): - if graph_params.is_mla: - update_decode_attention_mla_params( - update_stream, forward_meta, param, handle, event - ) - else: - update_decode_attention_params( - update_stream, forward_meta, param, handle, event - ) - - -def update_decode_attention_params(update_stream, forward_meta, param, handle, event): - ( - query, - key_cache, - value_cache, - num_kv_heads, - num_heads, - scale, - block_table, - kv_seq_len, - output, - ) = param - kv_seq_len = forward_meta.input_buffers["kv_seqlens"] - with torch.npu.stream(update_stream): - torch.npu.graph_task_update_begin(update_stream, handle) - torch.ops.atb._npu_paged_attention( - query=query, - key_cache=key_cache, - value_cache=value_cache, - num_kv_heads=num_kv_heads, - num_heads=num_heads, - scale_value=scale, - block_table=block_table, - context_lens=kv_seq_len, - out=output, - ) - torch.npu.graph_task_update_end(update_stream) - event.record(update_stream) - - -def update_decode_attention_mla_params( - update_stream, forward_meta, param, handle, event -): - ( - query, - key_cache, - num_kv_heads, - num_q_heads, - scale_value, - block_table, - kv_seq_len, - mla_vheadsize, - attn_output, - ) = param - kv_seq_len = forward_meta.input_buffers["kv_seqlens"] - with torch.npu.stream(update_stream): - torch.npu.graph_task_update_begin(update_stream, handle) - torch.ops.atb._npu_paged_attention_mla( - query=query, - key_cache=key_cache, - num_kv_heads=num_kv_heads, - num_heads=num_q_heads, - scale_value=scale_value, - block_table=block_table, - context_lens=kv_seq_len, - mla_vheadsize=mla_vheadsize, - out=attn_output, - ) - torch.npu.graph_task_update_end(update_stream) - event.record(update_stream) diff --git a/dlinfer/graph/dicp/dynamo_bridge/torch_version.py b/dlinfer/graph/dicp/dynamo_bridge/torch_version.py index 80687152..9f2ee836 100644 --- a/dlinfer/graph/dicp/dynamo_bridge/torch_version.py +++ b/dlinfer/graph/dicp/dynamo_bridge/torch_version.py @@ -11,6 +11,7 @@ is_torch_260 = False is_torch_271 = False is_torch_280 = False +is_torch_290 = False if torch_version.startswith("2.0"): @@ -29,6 +30,8 @@ is_torch_271 = True elif torch_version.startswith("2.8.0"): is_torch_280 = True +elif torch_version.startswith("2.9.0"): + is_torch_290 = True else: raise ValueError(f"unsupported dicp torch version: {torch.__version__}") diff --git a/dlinfer/ops/llm.py b/dlinfer/ops/llm.py index 0f9d9662..af2d8db1 100644 --- a/dlinfer/ops/llm.py +++ b/dlinfer/ops/llm.py @@ -419,6 +419,7 @@ def paged_prefill_attention( kv_scales, kv_zeros, quant_bits, + head_size_v, ) diff --git a/dlinfer/vendor/ascend/__init__.py b/dlinfer/vendor/ascend/__init__.py index fa256f0d..fd545b74 100644 --- a/dlinfer/vendor/ascend/__init__.py +++ b/dlinfer/vendor/ascend/__init__.py @@ -1,5 +1,6 @@ # Copyright (c) 2024, DeepLink. All rights reserved. -from pathlib import Path +from .version import ensure_ascend_runtime -import torch -from . import pytorch_patch, torch_npu_ops +ensure_ascend_runtime() + +from . import pytorch_patch, torch_npu_ops # noqa: E402,F401 diff --git a/dlinfer/vendor/ascend/attention.py b/dlinfer/vendor/ascend/attention.py index b4f8429e..dd342448 100644 --- a/dlinfer/vendor/ascend/attention.py +++ b/dlinfer/vendor/ascend/attention.py @@ -1,11 +1,7 @@ import math import torch +import torch_npu from dlinfer.utils.type_annotation import Tensor, Optional -from dlinfer.framework.lmdeploy_ext.cudagraph.ascend_cudagraph import ( - AscendGraphRunner, - get_graph_params, - aclgraph_use_torch_npu_update, -) def decode_attention( @@ -22,66 +18,29 @@ def decode_attention( softmax_scale: float, attn_output: Tensor, ): - if AscendGraphRunner.capturing and not aclgraph_use_torch_npu_update(): - graph_params = get_graph_params() - num_tokens = query.shape[0] - stream = torch.npu.current_stream() - event = torch.npu.ExternalEvent() - event.wait(stream) - event.reset(stream) - graph_params.events[num_tokens].append(event) - graph_params.attn_params[num_tokens].append( - ( - query, - key_cache, - value_cache, - num_kv_heads, - num_q_heads, - scale_value, - block_table, - kv_seq_len, - attn_output, - ) - ) - graph_params.is_mla = False - torch.npu.graph_task_group_begin(stream) - torch.ops.atb._npu_paged_attention( - query=query, - key_cache=key_cache, - value_cache=value_cache, - num_kv_heads=num_kv_heads, - num_heads=num_q_heads, - scale_value=scale_value, - block_table=block_table, - context_lens=kv_seq_len, - out=attn_output, - ) - handle = torch.npu.graph_task_group_end(stream) - graph_params.handles[num_tokens].append(handle) - else: - bs, _, dim = query.shape - block_num = key_cache.size(0) - query = query.contiguous() - attn_output = attn_output.contiguous() - key_cache = key_cache.view(block_num, block_size, -1) - value_cache = value_cache.view(block_num, block_size, -1) - scale_value = softmax_scale if softmax_scale else 1.0 / math.sqrt(dim) + _, _, dim = query.shape + block_num = key_cache.size(0) + query = query.contiguous() + attn_output = attn_output.contiguous() + key_cache = key_cache.view(block_num, block_size, -1) + value_cache = value_cache.view(block_num, block_size, -1) + scale_value = softmax_scale if softmax_scale else 1.0 / math.sqrt(dim) - attn_output, _ = torch.ops.npu.npu_fused_infer_attention_score( - query=query, - key=key_cache, - value=value_cache, - atten_mask=None, - block_table=block_table, - input_layout="TND", - block_size=block_size, - actual_seq_lengths=q_seq_len, - actual_seq_lengths_kv=kv_seq_len, - num_key_value_heads=num_kv_heads, - num_heads=num_q_heads, - scale=scale_value, - sparse_mode=0, - ) + attn_output, _ = torch.ops.npu.npu_fused_infer_attention_score( + query=query, + key=key_cache, + value=value_cache, + atten_mask=None, + block_table=block_table, + input_layout="TND", + block_size=block_size, + actual_seq_lengths=q_seq_len, + actual_seq_lengths_kv=kv_seq_len, + num_key_value_heads=num_kv_heads, + num_heads=num_q_heads, + scale=scale_value, + sparse_mode=0, + ) return attn_output @@ -96,52 +55,38 @@ def decode_attention_mla( mla_vheadsize: int, attn_output: Tensor, ): - if AscendGraphRunner.capturing: - graph_params = get_graph_params() - num_tokens = query.shape[0] - stream = torch.npu.current_stream() - event = torch.npu.ExternalEvent() - event.wait(stream) - event.reset(stream) - graph_params.events[num_tokens].append(event) - graph_params.attn_params[num_tokens].append( - ( - query, - key_cache, - num_kv_heads, - num_q_heads, - scale_value, - block_table, - kv_seq_len, - mla_vheadsize, - attn_output, - ) - ) - graph_params.is_mla = True - torch.npu.graph_task_group_begin(stream) - torch.ops.atb._npu_paged_attention_mla( - query=query, - key_cache=key_cache, - num_kv_heads=num_kv_heads, - num_heads=num_q_heads, - scale_value=scale_value, - block_table=block_table, - context_lens=kv_seq_len, - mla_vheadsize=mla_vheadsize, - out=attn_output, - ) - handle = torch.npu.graph_task_group_end(stream) - graph_params.handles[num_tokens].append(handle) - else: - torch.ops.atb._npu_paged_attention_mla( - query=query, - key_cache=key_cache, - num_kv_heads=num_kv_heads, - num_heads=num_q_heads, - scale_value=scale_value, - block_table=block_table, - context_lens=kv_seq_len, - mla_vheadsize=mla_vheadsize, - out=attn_output, - ) + num_tokens = query.shape[0] + _, block_size = key_cache.shape[:2] + + q_nope = ( + query[..., :mla_vheadsize] + .view(num_tokens, num_q_heads, 1, mla_vheadsize) + .contiguous() + ) + q_rope = query[..., mla_vheadsize:].view(num_tokens, num_q_heads, 1, -1) + + # FIA v2 expects paged KV cache in [block, kv_head, block_size, dim]. + key_cache = key_cache.permute(0, 2, 1, 3) + k_nope = key_cache[..., :mla_vheadsize] + k_rope = key_cache[..., mla_vheadsize:] + + fai_output, _ = torch_npu.npu_fused_infer_attention_score_v2( + q_nope, + k_nope, + k_nope, + query_rope=q_rope, + key_rope=k_rope, + num_query_heads=num_q_heads, + num_key_value_heads=num_kv_heads, + input_layout="BNSD_NBSD", + atten_mask=None, + sparse_mode=0, + softmax_scale=scale_value, + block_table=block_table, + block_size=block_size, + actual_seq_qlen=None, + actual_seq_kvlen=kv_seq_len, + ) + + attn_output.copy_(fai_output.squeeze(2).transpose(0, 1)) return attn_output diff --git a/dlinfer/vendor/ascend/torch_npu_ops.py b/dlinfer/vendor/ascend/torch_npu_ops.py index c8d7ef09..fa83b8c0 100644 --- a/dlinfer/vendor/ascend/torch_npu_ops.py +++ b/dlinfer/vendor/ascend/torch_npu_ops.py @@ -2,6 +2,7 @@ import math import torch import torch.distributed as dist +import torch_npu from typing import List from dlinfer.vendor import vendor_ops_registry @@ -197,43 +198,66 @@ def prefill_attention( else: # Handle qwenvl vision part flash-attention q_seq_len = get_cpu_seq_len(q_seq_len) - torch.ops.atb._npu_flash_attention_unpad( + is_tnd = ( + query.dim() == 3 + and query.shape[-2] == num_q_heads + and key.shape[-2] == num_kv_heads + ) + input_layout = "TND" if is_tnd else "BSH" + actual_seq_lengths = q_seq_len.cumsum(dim=0) if is_tnd else None + fia_kwargs = {} + if is_tnd and query.shape[-1] > value.shape[-1]: + nope_dim = value.shape[-1] + fia_kwargs["query_rope"] = query[..., nope_dim:].contiguous() + fia_kwargs["key_rope"] = key[..., nope_dim:].contiguous() + query = query[..., :nope_dim].contiguous() + key = key[..., :nope_dim].contiguous() + output, _ = torch_npu.npu_fused_infer_attention_score( query=query, key=key, value=value, - seq_len=q_seq_len, - scale_value=scale_value, + input_layout=input_layout, + actual_seq_lengths=actual_seq_lengths, + actual_seq_lengths_kv=actual_seq_lengths, + scale=scale_value, num_heads=num_q_heads, - num_kv_heads=num_kv_heads, - out=attn_output, + num_key_value_heads=num_kv_heads, + sparse_mode=0, + **fia_kwargs, ) + attn_output.copy_(output) return attn_output if SocVersion.is_Ascend910(): - torch.ops.atb._npu_flash_attention( + q_seq_len = get_cpu_seq_len(q_seq_len) + actual_seq_lengths = q_seq_len.cumsum(dim=0) + + # The backend supplies the fixed split-fuse causal mask required by + # sparse mode 3 for both standard attention and MLA. + fia_kwargs = {} + if query.shape[-1] > value.shape[-1]: + # MLA concatenates the NOPE and ROPE parts in Q/K. FIA accepts + # large MLA head dimensions only when the ROPE part is separate. + nope_dim = value.shape[-1] + fia_kwargs["query_rope"] = query[..., nope_dim:].contiguous() + fia_kwargs["key_rope"] = key[..., nope_dim:].contiguous() + query = query[..., :nope_dim].contiguous() + key = key[..., :nope_dim].contiguous() + + output, _ = torch_npu.npu_fused_infer_attention_score( query=query, key=key, value=value, - mask=mask, - seq_len=q_seq_len, - scale_value=scale_value, - num_heads=num_q_heads, - num_kv_heads=num_kv_heads, - out=attn_output, - ) - elif SocVersion.is_Ascend310P(): - # Used for Qwen2.5-VL model vision block - query = query.unsqueeze(0) - key = key.unsqueeze(0) - value = value.unsqueeze(0) - attn_output[:] = torch.ops.npu.npu_prompt_flash_attention( - query, - key, - value, + atten_mask=mask, + input_layout="TND", + actual_seq_lengths=actual_seq_lengths, + actual_seq_lengths_kv=actual_seq_lengths, + scale=scale_value, num_heads=num_q_heads, num_key_value_heads=num_kv_heads, - input_layout="BSND", - scale_value=scale_value, + sparse_mode=3, + **fia_kwargs, ) + attn_output.copy_(output) else: raise ValueError( f"dlinfer doesn't support {SocVersion.device_name()} device currently." @@ -448,19 +472,75 @@ def paged_prefill_attention( kv_scales: Optional[Tensor], kv_zeros: Optional[Tensor], quant_bits: Optional[int], + head_size_v: Optional[int] = None, ) -> Tensor: if alibi_slopes is not None: raise RuntimeError( "paged_decode_attention does not " "support alibi_slopes yet" ) + if isinstance(block_table, torch.Tensor) and block_table.dtype != torch.int32: + block_table = block_table.to(torch.int32) + scale_value = softmax_scale if softmax_scale else 1.0 / math.sqrt(query.shape[-1]) query = query.contiguous() + + # lmdeploy's DeepSeek MLA path absorbs W_UK into Q. Its paged cache is + # therefore [latent K/V, RoPE K], while value_cache is a view of the + # latent part. Follow vllm-ascend's multi-token paged attention path: + # feed the latent and RoPE components separately to FIA v2, and use the + # split-fuse causal mask with sparse mode 3. + is_mla = key_cache.shape[-1] != value_cache.shape[-1] + if is_mla: + num_tokens = query.shape[0] + mla_vheadsize = head_size_v or value_cache.shape[-1] + if query.shape[-1] <= mla_vheadsize: + raise RuntimeError( + "MLA paged prefill expects query to contain both latent and RoPE parts" + ) + + q_nope = query[..., :mla_vheadsize].contiguous() + q_rope = query[..., mla_vheadsize:].contiguous() + + # FIA v2 expects paged KV cache in + # [block, kv_head, block_size, dim]. + key_cache = key_cache.permute(0, 2, 1, 3) + value_cache = value_cache.permute(0, 2, 1, 3) + k_nope = key_cache[..., :mla_vheadsize] + k_rope = key_cache[..., mla_vheadsize:] + v_nope = value_cache[..., :mla_vheadsize] + mask = attn_mask[0] if len(attn_mask) else None + + output, _ = torch_npu.npu_fused_infer_attention_score_v2( + q_nope, + k_nope, + v_nope, + query_rope=q_rope, + key_rope=k_rope, + num_query_heads=num_q_heads, + num_key_value_heads=num_kv_heads, + input_layout="TND_NTD", + atten_mask=mask, + sparse_mode=3, + softmax_scale=scale_value, + block_table=block_table, + block_size=block_size, + actual_seq_qlen=q_seq_len, + actual_seq_kvlen=kv_seq_len, + ) + + # TND_NTD returns [num_heads, num_tokens, value_head_size]. + output = output[:, :num_tokens].transpose(0, 1) + if attn_output is not None: + attn_output.copy_(output) + return attn_output + return output + block_num = key_cache.size(0) key_cache = key_cache.view(block_num, block_size, -1) value_cache = value_cache.view(block_num, block_size, -1) - attn_output, _ = torch.ops.npu.npu_fused_infer_attention_score( + output, _ = torch.ops.npu.npu_fused_infer_attention_score( query=query, key=key_cache, value=value_cache, @@ -476,7 +556,10 @@ def paged_prefill_attention( sparse_mode=3, ) - return attn_output + if attn_output is not None: + attn_output.copy_(output) + return attn_output + return output @register_ops(vendor_ops_registry) diff --git a/dlinfer/vendor/ascend/version.py b/dlinfer/vendor/ascend/version.py new file mode 100644 index 00000000..ad5d3a18 --- /dev/null +++ b/dlinfer/vendor/ascend/version.py @@ -0,0 +1,64 @@ +# Copyright (c) 2026, DeepLink. All rights reserved. + +from typing import Optional + +from packaging.version import InvalidVersion, Version +import torch +import torch_npu + +MIN_TORCH_VERSION = Version("2.8.0") +MIN_TORCH_NPU_VERSION = Version("2.8.0.post1") + + +def _parse_version(raw_version: str, package_name: str) -> Version: + try: + return Version(raw_version) + except InvalidVersion as exc: + raise RuntimeError( + f"Invalid {package_name} version {raw_version!r}; DLINFER Ascend requires " + f"torch>={MIN_TORCH_VERSION} and torch-npu>={MIN_TORCH_NPU_VERSION}." + ) from exc + + +def ensure_ascend_runtime( + torch_version: Optional[str] = None, + torch_npu_version: Optional[str] = None, + check_graph_api: bool = True, +) -> tuple[Version, Version]: + """Validate the only supported Ascend graph-update runtime path.""" + parsed_torch = _parse_version( + torch.__version__ if torch_version is None else torch_version, "torch" + ) + parsed_torch_npu = _parse_version( + torch_npu.__version__ if torch_npu_version is None else torch_npu_version, + "torch-npu", + ) + + if parsed_torch < MIN_TORCH_VERSION: + raise RuntimeError( + f"Unsupported torch version {parsed_torch}; DLINFER Ascend requires " + f"torch>={MIN_TORCH_VERSION}. The legacy ATB graph-task update path " + "has been removed." + ) + if parsed_torch_npu < MIN_TORCH_NPU_VERSION: + raise RuntimeError( + f"Unsupported torch-npu version {parsed_torch_npu}; DLINFER Ascend " + f"requires torch-npu>={MIN_TORCH_NPU_VERSION}. The legacy ATB " + "graph-task update path has been removed." + ) + if parsed_torch.release[:2] != parsed_torch_npu.release[:2]: + raise RuntimeError( + "torch and torch-npu must use the same major.minor release for the " + f"Ascend backend, but got torch {parsed_torch} and torch-npu " + f"{parsed_torch_npu}." + ) + + graph_cls = getattr(getattr(torch, "npu", None), "NPUGraph", None) + if check_graph_api and not callable(getattr(graph_cls, "update", None)): + raise RuntimeError( + "torch.npu.NPUGraph.update is unavailable. DLINFER Ascend graph mode " + f"requires torch>={MIN_TORCH_VERSION} and " + f"torch-npu>={MIN_TORCH_NPU_VERSION}." + ) + + return parsed_torch, parsed_torch_npu diff --git a/dlinfer/vendor/camb/camb_ops.py b/dlinfer/vendor/camb/camb_ops.py index 11b5a1a3..874d388c 100644 --- a/dlinfer/vendor/camb/camb_ops.py +++ b/dlinfer/vendor/camb/camb_ops.py @@ -278,6 +278,7 @@ def paged_prefill_attention( kv_scales: Optional[Tensor], kv_zeros: Optional[Tensor], quant_bits: Optional[int], + head_size_v: Optional[int] = None, ) -> Tensor: if softmax_scale is None: softmax_scale = float(1 / math.sqrt(query.size(-1))) diff --git a/dlinfer/vendor/maca/maca_ops.py b/dlinfer/vendor/maca/maca_ops.py index a6426b52..e86e0cc5 100644 --- a/dlinfer/vendor/maca/maca_ops.py +++ b/dlinfer/vendor/maca/maca_ops.py @@ -294,6 +294,7 @@ def paged_prefill_attention( kv_scales: Optional[Tensor], kv_zeros: Optional[Tensor], quant_bits: Optional[int], + head_size_v: Optional[int] = None, ) -> Tensor: if softmax_scale is None: softmax_scale = float(1 / math.sqrt(query.size(-1))) diff --git a/docs/ascend_cudagraph_rl_lifecycle.md b/docs/ascend_cudagraph_rl_lifecycle.md new file mode 100644 index 00000000..73db6b81 --- /dev/null +++ b/docs/ascend_cudagraph_rl_lifecycle.md @@ -0,0 +1,331 @@ +# Ascend CUDAGraph 全局状态与 RL Re-capture 生命周期 + +本文说明 DLINFER Ascend CUDAGraph 实现中以下两个全局变量的历史用途、生命周期,以及它们与 RL rollout re-capture 的关系: + +```python +_graph_params: Optional[GraphParams] = None +_graph_capture_sizes: set[int] = None +``` + +核心结论如下: + +- `_graph_params` 历史上保存旧版 Attention graph-task 更新所需的图资源,必须在 rollout reset 时清理。 +- `_graph_capture_sizes` 当前只有赋值和清空,没有任何读取位置,因此没有实际控制作用。 +- RL re-capture 修复真正依赖的是 + `_get_capture_batch_size_impl.cache_clear()`,而不是 + `_graph_capture_sizes` 变量本身。 +- “capture size 需要跨 rollout 保存”指 capture-size 策略或配置,不是旧图的 handle、event 和 Tensor 引用。 + +> 实现状态(2026-08-17):本文第 6 节的重构已经落地。旧 ATB +> graph-task 路径以及 `_graph_params`、`_graph_capture_sizes` 已删除; +> capture-size 计算已成为纯函数,显式配置由 `CacheConfig` 跨 rollout 保存。 + +## 1. `_graph_params` 的历史用途 + +`_graph_params` 最早在提交 +[`5c474737`](https://github.com/DeepLink-org/dlinfer/commit/5c4747371b6bd3f39a71656c74e99e3429f259d1) +中引入,用于支持旧版 ATB Attention 的 graph-task 动态参数更新: + +```python +@dataclass +class GraphParams: + events: dict[int, list[torch.npu.ExternalEvent]] + workspaces: dict[int, torch.Tensor] + handles: dict[int, list[NPUTaskGroupHandle]] + attn_params: dict[int, list[tuple]] + is_mla: bool +``` + +其中的字典以 capture size 为 key。例如捕获 batch size 8 的图时,对应资源存放在: + +```python +events[8] +handles[8] +attn_params[8] +``` + +Attention capture 期间,每一层会注册以下信息: + +- Attention 输入 Tensor; +- KV cache Tensor; +- `block_table`; +- `kv_seq_len`; +- 输出 Tensor; +- graph-task handle; +- 同步 event。 + +旧版 replay 时,Graph Runner 根据当前 graph size 找到这些对象: + +```python +graph_params.attn_params[runtime_size] +graph_params.handles[runtime_size] +graph_params.events[runtime_size] +``` + +随后通过低层 graph-task API 更新 Attention 的运行时参数: + +```python +torch.npu.graph_task_update_begin(...) +torch.ops.atb._npu_paged_attention(...) +torch.npu.graph_task_update_end(...) +``` + +因此,`_graph_params` 本质上是旧 ATB Attention kernel 与 Graph +Runner 之间的全局 side channel。它不是普通配置元信息,而是持有真实图资源和 +Tensor 引用的对象。 + +## 2. 为什么 rollout 之间必须清理 `_graph_params` + +RL rollout 的典型生命周期如下: + +```text +rollout N 推理 + -> sleep / 更新权重 + -> reset graph + -> wakeup + -> rollout N+1 重新 capture +``` + +模型 sleep 或更新权重后: + +- 原权重 Tensor 地址可能发生变化; +- KV cache 可能被释放或重新分配; +- captured graph 中保存的输入输出地址可能已经失效; +- ATB task handle 和 event 属于旧图; +- `attn_params` 中保存的 Tensor 引用会阻止显存释放。 + +所以 `_graph_params` 不能跨 rollout 复用。 + +提交 +[`d0f60279`](https://github.com/DeepLink-org/dlinfer/commit/d0f60279a684de7dd37f0cab59d5065747faf598) +为此增加了 `clear_graph_params()`: + +```python +attn_params.clear() +handles.clear() +events.clear() +workspaces.clear() +_graph_params = None +``` + +如果错误地保留 `_graph_params`,可能导致: + +- 显存泄漏; +- 使用旧 KV cache 地址; +- 使用已失效的权重地址; +- graph replay 卡住; +- 权重更新后仍得到旧结果。 + +从生命周期上可以把状态分成两类: + +```text +应该跨 rollout 保存: + capture-size 策略、max_batches、模型类型 + +不应该跨 rollout 保存: + NPUGraph、task handle、event、Tensor 引用、output buffer +``` + +## 3. RL re-capture 问题的根因 + +历史实现把 capture-size 计算和 `_graph_params` 初始化放进了同一个带缓存函数: + +```python +@functools.lru_cache +def _get_capture_batch_size_impl(max_batches): + ... + set_graph_params(set(ret)) + return ret +``` + +这里混合了两种不同职责: + +1. 计算 capture sizes; +2. 初始化 `_graph_params`。 + +第一次 rollout 使用 `max_batches=256` 时: + +```text +_get_capture_batch_size_impl(256) + -> 函数体执行 + -> set_graph_params(...) + -> lru_cache 保存返回值 +``` + +sleep/reset 时: + +```text +clear_graph_params() + -> _graph_params = None +``` + +第二次 rollout 仍使用 `max_batches=256` 时: + +```text +_get_capture_batch_size_impl(256) + -> 命中 lru_cache + -> 直接返回 capture-size 列表 + -> 函数体不执行 + -> set_graph_params() 没有再次调用 + -> _graph_params 仍是 None +``` + +接下来 Attention capture 如果执行: + +```python +get_graph_params().events[...] +``` + +就会访问 `None`。这才是 re-capture 问题的根因。 + +## 4. `_graph_capture_sizes` 的实际作用 + +`_graph_capture_sizes` 由 +[`PR #319: fix re-capture in RL`](https://github.com/DeepLink-org/dlinfer/pull/319) +引入: + +```python +_graph_capture_sizes: set[int] = None +``` + +初始化时进行赋值: + +```python +_graph_capture_sizes = aclgraph_capture_sizes +``` + +清理时执行: + +```python +_graph_capture_sizes = None +_get_capture_batch_size_impl.cache_clear() +``` + +检查该 PR 的提交快照、当前开发分支和 `origin/main` 后可以确认,`_graph_capture_sizes` 只有以下操作: + +- 声明; +- `global` 引用; +- 赋值; +- 清空。 + +代码中没有读取它的 getter,也没有用它重新初始化 `GraphParams`。因此,按照当前实现, +`_graph_capture_sizes` 是一个冗余的 bookkeeping 变量,对 re-capture +没有功能性贡献。 + +PR #319 真正修复问题的是: + +```python +_get_capture_batch_size_impl.cache_clear() +``` + +它保证下一次 rollout 即使使用相同的 `max_batches`,函数体也会重新执行,并再次调用 `set_graph_params()`。 + +该 PR 没有正文和讨论,只有 8 行修改。因此无法从历史材料证明作者计划在其他地方读取 `_graph_capture_sizes`;从现有代码只能确认它没有参与任何控制逻辑。 + +## 5. 如何理解“capture size 要跨 rollout 保存” + +这里的 capture size 应当理解为 engine 需要捕获哪些 batch size,例如: + +```text +[1, 2, 4, 8, 16, 32] +``` + +这是配置或 shape policy,应该跨 rollout 保留,因为它决定 wakeup 后需要重新捕获哪些图。 + +但它不等同于 `_graph_capture_sizes` 这个全局变量。更合理的持久化来源是: + +```python +CacheConfig.max_batches +CacheConfig.cudagraph_capture_batch_sizes +``` + +LMDeploy 后来在提交 +[`4f25485e`](https://github.com/InternLM/lmdeploy/commit/4f25485e218d5e0938240da1d4fccbb426573a89) +中正式把 capture sizes 放进了 `CacheConfig`: + +```python +cudagraph_capture_batch_sizes: list[int] | None +``` + +合理的状态归属如下: + +```text +CacheConfig 跨 rollout 保留 + `-- cudagraph_capture_batch_sizes + +AscendGraphRunner rollout 之间 reset + `-- _runner_map + `-- AscendSingleGraphRunner + `-- NPUGraph + +旧版全局 _graph_params rollout 之间销毁 + |-- handles + |-- events + `-- Tensor refs +``` + +重构后的 Ascend DLINFER `get_capture_batch_sizes()` 会优先读取 LMDeploy 的 +`CacheConfig.cudagraph_capture_batch_sizes`;未显式配置时,才调用 Ascend 的 +`_get_capture_batch_size_impl()` 生成默认尺寸。 + +## 6. 本次重构结果 + +本次重构没有只做以下机械删除: + +```text +删除 _graph_params +删除 _graph_capture_sizes +删除 cache_clear() +``` + +同时完成了以下生命周期迁移: + +1. capture-size 计算已经成为纯函数: + + ```python + @functools.lru_cache + def _get_capture_batch_size_impl(max_batches): + ... + return ret + ``` + + 其中不能再调用 `set_graph_params()`。 + +2. capture-size 策略保存在 `CacheConfig` 中,可以跨 rollout 保留。 +3. `op_backend.py` 根据 K/V head dim 显式判断并向 Graph Runner 传递 `is_mla`。 +4. Dense decode 固定使用 FIA,MLA decode 和 paged prefill 固定使用 FIA v2。 +5. Graph replay 固定使用 `NPUGraph.update()`;旧 ATB graph-task Attention 路径已删除。 +6. `_graph_params` 和从未被读取的 `_graph_capture_sizes` 已整体删除。 +7. `AscendGraphRunner.reset()` 先调用基类 reset,清空 rollout 局部的 + `padding_batch_size`,但不清除 `CacheConfig` 中的 capture sizes。 +8. 增加 capture-size 纯函数、reset 配置保留、FIA v2 graph replay 和版本门槛测试。 + +完整 RL 场景仍建议执行以下端到端回归: + +```text +capture -> inference +reset/sleep +wakeup +使用相同 max_batches 再次 capture +inference +重复两到三轮 +``` + +测试需要验证: + +- 不出现 `_graph_params is None`; +- 不使用旧权重; +- 不持有旧 KV cache; +- capture sizes 与第一次一致; +- graph replay 精度正确; +- 多轮 sleep/wakeup 后没有显存持续增长。 + +## 7. 最终结论 + +最终可以删除 `_graph_params` 和 `_graph_capture_sizes`,但必须保留它们背后的两类语义: + +- `_graph_params` 背后的旧图资源生命周期:新路径中由每个 `NPUGraph` 自己管理,并在 reset 时释放。 +- `_graph_capture_sizes` 名字所暗示的 capture-size 策略:迁移到持久的 + `CacheConfig`,不能随着 rollout 丢失。 + +换句话说,需要保留的是生命周期语义和配置来源,而不是这两个模块级全局变量本身。 diff --git a/requirements/ascend/torch.txt b/requirements/ascend/torch.txt index f99c0aca..3a7fe4e2 100644 --- a/requirements/ascend/torch.txt +++ b/requirements/ascend/torch.txt @@ -1,6 +1,7 @@ # Please install one of the supported versions manually -torch>=2.3.1,<2.10.0 -torch-npu>=2.3.1,<2.10.0 -torchvision>=0.18.1,<0.25.0 +torch>=2.8.0,<2.10.0 +torch-npu>=2.8.0.post1,<2.10.0 +torchvision>=0.23.0,<0.25.0 importlib-metadata +packaging pyyaml diff --git a/tests/test_ascend_attention_precision.py b/tests/test_ascend_attention_precision.py new file mode 100644 index 00000000..abb43f71 --- /dev/null +++ b/tests/test_ascend_attention_precision.py @@ -0,0 +1,358 @@ +# Copyright (c) 2026, DeepLink. All rights reserved. + +import math + +import pytest +import torch + +torch_npu = pytest.importorskip("torch_npu") + +pytestmark = pytest.mark.lmdeploy + +if not torch.npu.is_available(): + pytest.skip("Ascend NPU is required", allow_module_level=True) + +from dlinfer.ops import paged_prefill_attention +from dlinfer.vendor.ascend.attention import decode_attention_mla +from dlinfer.vendor.ascend.torch_npu_ops import prefill_attention + +DTYPE = torch.bfloat16 +DEVICE = torch.device("npu") +FAI_CAUSAL_MASK_SIZE = 2048 +NUM_Q_HEADS = 16 +NUM_KV_HEADS = 1 +STANDARD_HEAD_DIM = 128 +STANDARD_SOFTMAX_SCALE = 1.0 / math.sqrt(STANDARD_HEAD_DIM) +DENSE_NOPE_HEAD_DIM = 128 +DENSE_ROPE_HEAD_DIM = 64 +DENSE_QK_HEAD_DIM = DENSE_NOPE_HEAD_DIM + DENSE_ROPE_HEAD_DIM +DENSE_V_HEAD_DIM = DENSE_NOPE_HEAD_DIM +DENSE_SOFTMAX_SCALE = 1.0 / math.sqrt(DENSE_QK_HEAD_DIM) +MLA_NOPE_HEAD_DIM = 512 +MLA_ROPE_HEAD_DIM = 64 +MLA_QK_HEAD_DIM = MLA_NOPE_HEAD_DIM + MLA_ROPE_HEAD_DIM +MLA_V_HEAD_DIM = MLA_NOPE_HEAD_DIM +# DeepSeekV2 keeps the scale of its pre-absorption 128 + 64 query head. +MLA_SOFTMAX_SCALE = 1.0 / math.sqrt(128 + MLA_ROPE_HEAD_DIM) + + +@pytest.fixture(scope="module") +def fai_causal_mask(): + """Build the fixed split-fuse mask expected by FAI sparse mode 3.""" + return torch.triu( + torch.ones( + FAI_CAUSAL_MASK_SIZE, + FAI_CAUSAL_MASK_SIZE, + dtype=torch.int8, + device=DEVICE, + ), + diagonal=1, + ) + + +def _randn(shape): + return torch.randn(shape, dtype=torch.float32).to(DTYPE) + + +def _repeat_kv(hidden_states, num_q_heads): + num_kv_heads = hidden_states.shape[1] + assert num_q_heads % num_kv_heads == 0 + return hidden_states.repeat_interleave(num_q_heads // num_kv_heads, dim=1) + + +def _torch_prefill_attention(query, key, value, seq_lens, softmax_scale): + outputs = [] + start = 0 + for seq_len in seq_lens: + end = start + seq_len + q = query[start:end].float().transpose(0, 1) + k = _repeat_kv(key[start:end], NUM_Q_HEADS).float().transpose(0, 1) + v = _repeat_kv(value[start:end], NUM_Q_HEADS).float().transpose(0, 1) + + scores = torch.matmul(q, k.transpose(-1, -2)) * softmax_scale + causal_mask = torch.triu( + torch.ones(seq_len, seq_len, dtype=torch.bool), diagonal=1 + ) + scores.masked_fill_(causal_mask, float("-inf")) + outputs.append(torch.matmul(torch.softmax(scores, dim=-1), v).transpose(0, 1)) + start = end + + return torch.cat(outputs) + + +def _torch_decode_attention(query, key_cache, block_table, kv_seq_lens): + outputs = [] + block_size = key_cache.shape[1] + for batch_idx, kv_seq_len in enumerate(kv_seq_lens): + num_blocks = math.ceil(kv_seq_len / block_size) + block_ids = block_table[batch_idx, :num_blocks].long() + key = key_cache[block_ids].flatten(0, 1)[:kv_seq_len] + value = key[..., :MLA_V_HEAD_DIM] + + q = query[batch_idx].float() + k = _repeat_kv(key, NUM_Q_HEADS).float().transpose(0, 1) + v = _repeat_kv(value, NUM_Q_HEADS).float().transpose(0, 1) + scores = torch.matmul(q.unsqueeze(1), k.transpose(-1, -2)) + scores = scores * MLA_SOFTMAX_SCALE + outputs.append(torch.matmul(torch.softmax(scores, dim=-1), v).squeeze(1)) + + return torch.stack(outputs) + + +def _torch_paged_prefill_attention( + query, key_cache, block_table, q_seq_lens, kv_seq_lens +): + outputs = [] + query_start = 0 + block_size = key_cache.shape[1] + for batch_idx, (q_seq_len, kv_seq_len) in enumerate(zip(q_seq_lens, kv_seq_lens)): + num_blocks = math.ceil(kv_seq_len / block_size) + block_ids = block_table[batch_idx, :num_blocks].long() + key = key_cache[block_ids].flatten(0, 1)[:kv_seq_len] + value = key[..., :MLA_V_HEAD_DIM] + query_end = query_start + q_seq_len + + q = query[query_start:query_end].float().transpose(0, 1) + k = _repeat_kv(key, NUM_Q_HEADS).float().transpose(0, 1) + v = _repeat_kv(value, NUM_Q_HEADS).float().transpose(0, 1) + scores = torch.matmul(q, k.transpose(-1, -2)) * MLA_SOFTMAX_SCALE + + history_len = kv_seq_len - q_seq_len + q_positions = history_len + torch.arange(q_seq_len) + kv_positions = torch.arange(kv_seq_len) + causal_mask = kv_positions.unsqueeze(0) > q_positions.unsqueeze(1) + scores.masked_fill_(causal_mask.unsqueeze(0), float("-inf")) + outputs.append(torch.matmul(torch.softmax(scores, dim=-1), v).transpose(0, 1)) + query_start = query_end + + return torch.cat(outputs) + + +def _assert_prefill_attention_matches_torch( + qk_head_dim, value_head_dim, softmax_scale, causal_mask +): + seq_lens = [112, 70, 31] + num_tokens = sum(seq_lens) + + query = _randn((num_tokens, NUM_Q_HEADS, qk_head_dim)) + key = _randn((num_tokens, NUM_KV_HEADS, qk_head_dim)) + value = _randn((num_tokens, NUM_KV_HEADS, value_head_dim)) + expected = _torch_prefill_attention(query, key, value, seq_lens, softmax_scale) + + query = query.to(DEVICE) + key = key.to(DEVICE) + value = value.to(DEVICE) + seq_lens_tensor = torch.tensor(seq_lens, dtype=torch.int32) + max_seq_len = max(seq_lens) + output = torch.empty( + (num_tokens, NUM_Q_HEADS, value_head_dim), dtype=DTYPE, device=DEVICE + ) + + actual = prefill_attention( + query=query, + key=key, + value=value, + q_start_loc=None, + q_seq_len=seq_lens_tensor, + max_q_seq_len=max_seq_len, + num_q_heads=NUM_Q_HEADS, + num_kv_heads=NUM_KV_HEADS, + attn_mask=[causal_mask], + softmax_scale=softmax_scale, + alibi_slopes=None, + attn_output=output, + ) + + assert actual.data_ptr() == output.data_ptr() + torch.testing.assert_close(actual.cpu().float(), expected, rtol=5e-3, atol=5e-3) + + +def test_prefill_attention_standard_matches_torch(fai_causal_mask): + torch.manual_seed(20260812) + _assert_prefill_attention_matches_torch( + qk_head_dim=STANDARD_HEAD_DIM, + value_head_dim=STANDARD_HEAD_DIM, + softmax_scale=STANDARD_SOFTMAX_SCALE, + causal_mask=fai_causal_mask, + ) + + +def test_prefill_attention_dense_matches_torch(fai_causal_mask): + torch.manual_seed(20260813) + _assert_prefill_attention_matches_torch( + qk_head_dim=DENSE_QK_HEAD_DIM, + value_head_dim=DENSE_V_HEAD_DIM, + softmax_scale=DENSE_SOFTMAX_SCALE, + causal_mask=fai_causal_mask, + ) + + +def test_prefill_attention_mla_matches_torch(fai_causal_mask): + torch.manual_seed(20260814) + _assert_prefill_attention_matches_torch( + qk_head_dim=MLA_QK_HEAD_DIM, + value_head_dim=MLA_V_HEAD_DIM, + softmax_scale=MLA_SOFTMAX_SCALE, + causal_mask=fai_causal_mask, + ) + + +def test_paged_prefill_attention_mla_matches_torch(fai_causal_mask): + torch.manual_seed(20260814) + q_seq_lens = [3, 2] + kv_seq_lens = [130, 77] + cumulative_q_seq_lens = [3, 5] + block_size = 128 + num_blocks = 4 + block_table = torch.tensor([[2, 0], [3, 1]], dtype=torch.int32) + + query = _randn((sum(q_seq_lens), NUM_Q_HEADS, MLA_QK_HEAD_DIM)) + key_cache = _randn((num_blocks, block_size, NUM_KV_HEADS, MLA_QK_HEAD_DIM)) + expected = _torch_paged_prefill_attention( + query, key_cache, block_table, q_seq_lens, kv_seq_lens + ) + + query = query.to(DEVICE) + key_cache = key_cache.to(DEVICE) + value_cache = key_cache[..., :MLA_V_HEAD_DIM] + output = torch.empty( + (sum(q_seq_lens), NUM_Q_HEADS, MLA_V_HEAD_DIM), + dtype=DTYPE, + device=DEVICE, + ) + + actual = paged_prefill_attention( + query=query, + key=query[:, :NUM_KV_HEADS], + value=query[:, :NUM_KV_HEADS, :MLA_V_HEAD_DIM], + key_cache=key_cache, + value_cache=value_cache, + block_table=block_table.to(DEVICE), + block_size=block_size, + q_start_loc=None, + q_seq_len=torch.tensor(cumulative_q_seq_lens, dtype=torch.int32), + kv_seq_len=torch.tensor(kv_seq_lens, dtype=torch.int32), + cu_seq_lens_kv=None, + max_q_seq_len=max(q_seq_lens), + max_kv_seq_len=max(kv_seq_lens), + num_q_heads=NUM_Q_HEADS, + num_kv_heads=NUM_KV_HEADS, + attn_mask=[fai_causal_mask], + softmax_scale=MLA_SOFTMAX_SCALE, + alibi_slopes=None, + attn_output=output, + kv_scales=None, + kv_zeros=None, + quant_bits=0, + head_size_v=MLA_V_HEAD_DIM, + ) + + assert actual.data_ptr() == output.data_ptr() + torch.testing.assert_close(actual.cpu().float(), expected, rtol=5e-3, atol=5e-3) + + +def test_paged_prefill_attention_mla_graph_replay(fai_causal_mask): + torch.manual_seed(20260814) + q_seq_lens = [2] + capture_kv_seq_lens = [9] + replay_kv_seq_lens = [7] + block_size = 128 + block_table = torch.tensor([[0]], dtype=torch.int32) + query = _randn((sum(q_seq_lens), NUM_Q_HEADS, MLA_QK_HEAD_DIM)) + key_cache = _randn((1, block_size, NUM_KV_HEADS, MLA_QK_HEAD_DIM)) + expected = _torch_paged_prefill_attention( + query, key_cache, block_table, q_seq_lens, replay_kv_seq_lens + ) + + query = query.to(DEVICE) + key_cache = key_cache.to(DEVICE) + output = torch.empty( + (sum(q_seq_lens), NUM_Q_HEADS, MLA_V_HEAD_DIM), + dtype=DTYPE, + device=DEVICE, + ) + kwargs = dict( + query=query, + key=query[:, :NUM_KV_HEADS], + value=query[:, :NUM_KV_HEADS, :MLA_V_HEAD_DIM], + key_cache=key_cache, + value_cache=key_cache[..., :MLA_V_HEAD_DIM], + block_table=block_table.to(DEVICE), + block_size=block_size, + q_start_loc=None, + q_seq_len=torch.tensor(q_seq_lens, dtype=torch.int32), + kv_seq_len=torch.tensor(capture_kv_seq_lens, dtype=torch.int32), + cu_seq_lens_kv=None, + max_q_seq_len=max(q_seq_lens), + max_kv_seq_len=max(capture_kv_seq_lens), + num_q_heads=NUM_Q_HEADS, + num_kv_heads=NUM_KV_HEADS, + attn_mask=[fai_causal_mask], + softmax_scale=MLA_SOFTMAX_SCALE, + alibi_slopes=None, + attn_output=output, + kv_scales=None, + kv_zeros=None, + quant_bits=0, + head_size_v=MLA_V_HEAD_DIM, + ) + + # Warm up allocations before capture, matching the model graph runner. + paged_prefill_attention(**kwargs) + torch.npu.synchronize() + + graph = torch.npu.NPUGraph() + capture_stream = torch.npu.Stream() + try: + with torch.npu.graph( + graph, + auto_dispatch_capture=True, + stream=capture_stream, + ): + actual = paged_prefill_attention(**kwargs) + + graph.replay() + graph.update(cpu_update_input=[{"actual_seq_kvlen": replay_kv_seq_lens}]) + torch.npu.synchronize() + + assert actual.data_ptr() == output.data_ptr() + torch.testing.assert_close(actual.cpu().float(), expected, rtol=5e-3, atol=5e-3) + finally: + graph.reset() + + +def test_decode_attention_mla_matches_torch(): + torch.manual_seed(20260813) + batch_size = 3 + block_size = 128 + num_blocks = 6 + kv_seq_lens = [130, 77, 17] + block_table = torch.tensor( + [[4, 1], [3, 0], [5, 2]], + dtype=torch.int32, + ) + + query = _randn((batch_size, NUM_Q_HEADS, MLA_QK_HEAD_DIM)) + key_cache = _randn((num_blocks, block_size, NUM_KV_HEADS, MLA_QK_HEAD_DIM)) + expected = _torch_decode_attention(query, key_cache, block_table, kv_seq_lens) + + query = query.to(DEVICE) + key_cache = key_cache.to(DEVICE) + output = torch.empty( + (batch_size, NUM_Q_HEADS, MLA_V_HEAD_DIM), dtype=DTYPE, device=DEVICE + ) + + actual = decode_attention_mla( + query=query, + key_cache=key_cache, + num_kv_heads=NUM_KV_HEADS, + num_q_heads=NUM_Q_HEADS, + scale_value=MLA_SOFTMAX_SCALE, + block_table=block_table.to(DEVICE), + kv_seq_len=torch.tensor(kv_seq_lens, dtype=torch.int32), + mla_vheadsize=MLA_V_HEAD_DIM, + attn_output=output, + ) + + assert actual.data_ptr() == output.data_ptr() + torch.testing.assert_close(actual.cpu().float(), expected, rtol=5e-3, atol=5e-3) diff --git a/tests/test_ascend_cudagraph_lifecycle.py b/tests/test_ascend_cudagraph_lifecycle.py new file mode 100644 index 00000000..a42a4df7 --- /dev/null +++ b/tests/test_ascend_cudagraph_lifecycle.py @@ -0,0 +1,44 @@ +# Copyright (c) 2026, DeepLink. All rights reserved. + +from types import SimpleNamespace + +import pytest +import torch + +pytest.importorskip("torch_npu") + +from dlinfer.framework.lmdeploy_ext.cudagraph.ascend_cudagraph import ( + AscendGraphRunner, + _get_capture_batch_size_impl, +) +from lmdeploy.pytorch.backends.graph_runner import GraphRunnerMeta + + +def test_capture_size_generation_is_pure(): + _get_capture_batch_size_impl.cache_clear() + first = _get_capture_batch_size_impl(33) + _get_capture_batch_size_impl.cache_clear() + second = _get_capture_batch_size_impl(33) + + assert first == second + assert first[-1] == 33 + + +def test_reset_preserves_configured_capture_sizes(monkeypatch): + runner = object.__new__(AscendGraphRunner) + runner._runner_meta = GraphRunnerMeta(padding_batch_size=16) + runner._runner_map = {} + runner.graph_pool_handle = object() + runner.cache_config = SimpleNamespace( + max_batches=32, + cudagraph_capture_batch_sizes=[1, 4, 16, 32], + ) + monkeypatch.setattr(torch.npu, "empty_cache", lambda: None) + + before_reset = runner.get_capture_batch_sizes().copy() + runner.reset() + after_reset = runner.get_capture_batch_sizes().copy() + + assert runner.get_meta().padding_batch_size is None + assert runner.graph_pool_handle is None + assert before_reset == after_reset == [1, 4, 16, 32]