# # Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. # """NPU-specific encoder ACL graph: params, runtime context, FIA replay updates, and manager.""" from __future__ import annotations from contextlib import contextmanager from dataclasses import dataclass, field from typing import Any import torch import torch_npu from vllm.logger import logger from vllm.platforms import current_platform from vllm.v1.worker.encoder_cudagraph import BudgetGraphMetadata, EncoderCudaGraphManager from vllm_ascend.utils import vllm_version_is, weak_ref_tensors # --------------------------------------------------------------------------- # Per–encoder-budget ACL graph bookkeeping (ViT FIA tasks) # --------------------------------------------------------------------------- @dataclass class EncoderGraphParams: """Mirrors :class:`vllm_ascend.compilation.acl_graph.GraphParams` but keyed by encoder token budget.""" # TODO: Fully support upstream dual-path encoder graph on Ascend. The # current FIA bookkeeping is keyed only by token_budget; dual-path models # need graph params separated by (path, token_budget). events: dict[int, list[torch.npu.ExternalEvent]] = field(default_factory=dict) workspaces: dict[int, torch.Tensor | None] = field(default_factory=dict) handles: dict[int, list[Any]] = field(default_factory=dict) # Flattened per-forward insertion order (one entry per ViT block invocation). attn_params: dict[int, list[tuple]] = field(default_factory=dict) _encoder_graph_params: EncoderGraphParams | None = None def set_encoder_graph_params(token_budgets: list[int]) -> None: global _encoder_graph_params budgets_sorted_unique = sorted(token_budgets) _encoder_graph_params = EncoderGraphParams( events={b: [] for b in budgets_sorted_unique}, workspaces={b: None for b in budgets_sorted_unique}, handles={b: [] for b in budgets_sorted_unique}, attn_params={b: [] for b in budgets_sorted_unique}, ) def get_encoder_graph_params() -> EncoderGraphParams | None: return _encoder_graph_params def update_encoder_graph_workspace(token_budget: int, workspace: torch.Tensor) -> None: if _encoder_graph_params is None: return _encoder_graph_params.workspaces[token_budget] = workspace # --------------------------------------------------------------------------- # Capture / replay runtime state (thread-local module singleton) # --------------------------------------------------------------------------- @dataclass class EncoderForwardContext: """Vision encoder NPUGraph runtime flags and host-side FIA arguments. Captured tensors stay on device; FIA ``graph_task_update`` needs Python ``list[int]`` lengths that are refreshed each replay from encoder metadata buffers on device (see RFC). """ token_budget: int | None = None capturing: bool = False cu_seqlens_cpu: torch.Tensor | None = None _context = EncoderForwardContext() def get_encoder_forward_context() -> EncoderForwardContext: return _context def _reset_encoder_forward_context() -> None: """Clear replay-time host length fields.""" _context.token_budget = None _context.capturing = False _context.cu_seqlens_cpu = None @contextmanager def set_encoder_forward_context( token_budget: int, capturing: bool, *, cu_seqlens_cpu: torch.Tensor | None = None, ): """Enter encoder graph replay (FIA host args): callers must pass lengths each time. On exit, replay host fields are **cleared** (not restored). Tensors must not be reused across replays without repopulating from the current batch buffers. """ _context.token_budget = token_budget _context.capturing = capturing _context.cu_seqlens_cpu = cu_seqlens_cpu try: yield _context finally: _reset_encoder_forward_context() # --------------------------------------------------------------------------- # FIA actual_seq_lengths (cu_seqlens -> list[int]) # --------------------------------------------------------------------------- def maybe_compute_actual_seq_lengths( cu_seqlens: torch.Tensor, num_query_tokens: int, num_kv_tokens: int, *, cudagraph_mm_encoder: bool = False, ) -> tuple[list[int], list[int]]: """Convert ``cu_seqlens`` to FIA host ``actual_seq_lengths``. Drops the leading-zero marker; with ``cudagraph_mm_encoder`` filters endpoints and aligns the terminal to ``num_query_tokens``; when Q≠KV, scales Q endpoints by ``num_kv_tokens // num_query_tokens`` for the KV list. """ flat = cu_seqlens.detach().cpu().view(-1).tolist() actual = flat[1:] if flat else flat if not cudagraph_mm_encoder: pass elif num_query_tokens <= 0: actual = [0] else: filtered: list[int] = [] for end in actual: if end <= 0: continue if end > num_query_tokens: break if not filtered or end > filtered[-1]: filtered.append(end) if not filtered or filtered[-1] != num_query_tokens: filtered.append(num_query_tokens) actual = filtered if num_kv_tokens == num_query_tokens: return actual, actual assert num_kv_tokens % num_query_tokens == 0 ratio = num_kv_tokens // num_query_tokens return actual, [end * ratio for end in actual] def update_encoder_graph_params( update_stream: torch.npu.Stream, token_budget: int, ) -> None: """Re-bind fused infer attention host tensors inside the encoder NPUGraph (parallel to LLM path). This deliberately bypasses :class:`AttentionBackend` — ViT attention is not registered there — but reuses the same ``graph_task_update_{begin,end}`` + ``ExternalEvent`` ordering pattern as :meth:`AscendAttentionBackendImpl.update_graph_params`. """ params = get_encoder_graph_params() if params is None or token_budget not in params.handles: return handles = params.handles[token_budget] events = params.events[token_budget] attn_blocks = params.attn_params[token_budget] workspace = params.workspaces.get(token_budget) if len(handles) != len(events) or len(handles) != len(attn_blocks): raise RuntimeError( "Encoder graph bookkeeping is inconsistent: " f"budget={token_budget} handles={len(handles)} " f"events={len(events)} attn_blocks={len(attn_blocks)}" ) with torch.npu.stream(update_stream): for handle, event, packed in zip(handles, events, attn_blocks): ( query, key, value, block_table, attn_mask, block_size, num_kv_heads, num_heads, scale, output, softmax_lse, ) = packed num_query_tokens = query.shape[0] num_kv_tokens = key.shape[0] cu_seqlens_cpu = get_encoder_forward_context().cu_seqlens_cpu actual_seq_lengths_q, actual_seq_lengths_kv = maybe_compute_actual_seq_lengths( cu_seqlens_cpu, num_query_tokens, num_kv_tokens, cudagraph_mm_encoder=True, ) torch.npu.graph_task_update_begin(update_stream, handle) torch_npu.npu_fused_infer_attention_score.out( query=query, key=key, value=value, atten_mask=attn_mask, block_table=block_table, input_layout="TND", block_size=block_size, actual_seq_lengths=actual_seq_lengths_q, actual_seq_lengths_kv=actual_seq_lengths_kv, num_key_value_heads=num_kv_heads, num_heads=num_heads, scale=scale, sparse_mode=0, workspace=workspace, out=[output, softmax_lse], ) torch.npu.graph_task_update_end(update_stream) event.record(update_stream) # --------------------------------------------------------------------------- # Encoder NPUGraph manager # --------------------------------------------------------------------------- class EncoderAclGraphManager(EncoderCudaGraphManager): """Hooks encoder capture/replay into Ascend FIA graph-task infrastructure.""" def __init__(self, *args: Any, **kwargs: Any) -> None: super().__init__(*args, **kwargs) self.graph_pool = current_platform.get_global_graph_pool() self.update_stream: torch.npu.Stream | None = None def capture(self, graph_pool: Any | None = None): encoder_graph_pool = graph_pool if graph_pool is not None else self.graph_pool self.graph_pool = encoder_graph_pool set_encoder_graph_params(self.token_budgets) super().capture(graph_pool=encoder_graph_pool) weak_ref_workspaces() def _capture_budget_graph(self, token_budget: int, path: str = "default"): logger.debug( "Capturing encoder aclgraph for budget=%d, max_batch_size=%d, max_frames_per_batch=%d", token_budget, self.max_batch_size, self.max_frames_per_batch, ) if vllm_version_is("0.23.0"): capture_inputs = self.model.prepare_encoder_cudagraph_capture_inputs( token_budget, self.max_batch_size, self.max_frames_per_batch, self.device, self.dtype, ) else: capture_inputs = self.model.prepare_encoder_cudagraph_capture_inputs( token_budget, self.max_batch_size, self.max_frames_per_batch, self.device, self.dtype, path, ) values = capture_inputs.values with torch.inference_mode(): if vllm_version_is("0.23.0"): output = self.model.encoder_cudagraph_forward(dict(values)) else: output = self.model.encoder_cudagraph_forward(dict(values), path=path) output_buffer = torch.empty_like(output) graph = torch.npu.NPUGraph() with ( set_encoder_forward_context(token_budget, True), torch.inference_mode(), torch.npu.graph(graph, self.graph_pool), ): if vllm_version_is("0.23.0"): output = self.model.encoder_cudagraph_forward(dict(values)) else: output = self.model.encoder_cudagraph_forward(dict(values), path=path) output_buffer.copy_(output) graph_meta = BudgetGraphMetadata( token_budget=token_budget, max_batch_size=self.max_batch_size, max_frames_per_batch=self.max_frames_per_batch, graph=graph, input_buffers=values, output_buffer=weak_ref_tensors(output_buffer), ) if vllm_version_is("0.23.0"): self.budget_graphs[token_budget] = graph_meta else: graph_set = self._get_graph_set(path) graph_set[token_budget] = graph_meta def _run_budget_graph( self, mm_kwargs: dict[str, Any], token_budget: int, path: str = "default", ) -> torch.Tensor | None: num_items = len(self._get_item_specs(mm_kwargs)) if vllm_version_is("0.23.0"): if token_budget not in self.budget_graphs: self.graph_misses += num_items return None graph_meta = self.budget_graphs[token_budget] else: graph_set = self._get_graph_set(path) if token_budget not in graph_set: self.graph_misses += num_items return None graph_meta = graph_set[token_budget] if vllm_version_is("0.23.0"): replay = self.model.prepare_encoder_cudagraph_replay_buffers( mm_kwargs, self.max_batch_size, self.max_frames_per_batch, ) buffer_items = ((key, graph_meta.input_buffers[key]) for key in self.config.buffer_keys) else: replay = self.model.prepare_encoder_cudagraph_replay_buffers( mm_kwargs, self.max_batch_size, self.max_frames_per_batch, path, ) buffer_items = graph_meta.input_buffers.items() for key, buf in buffer_items: src = replay.values.get(key) if src is None: continue if src.ndim == 0: buf.copy_(src) else: padding_logic = self.config.padding_logics.get(key, self._copy_padded_buffer) padding_logic(buf, src) cu_seqlens = graph_meta.input_buffers.get("cu_seqlens") cu_seqlens_cpu = None if cu_seqlens is None else cu_seqlens.cpu() update_stream = self.update_stream if update_stream is None: update_stream = torch.npu.Stream() graph_meta.graph.replay() with set_encoder_forward_context( token_budget, False, cu_seqlens_cpu=cu_seqlens_cpu, ): update_encoder_graph_params(update_stream, token_budget) self.graph_hits += num_items return graph_meta.output_buffer def weak_ref_workspaces() -> None: params = get_encoder_graph_params() if params is None: return for budget, ws in list(params.workspaces.items()): if ws is None: continue params.workspaces[budget] = weak_ref_tensors(ws)