Files
enginex-ascend-910-vllm/vllm_ascend/worker/encoder_acl_graph.py

407 lines
14 KiB
Python
Raw Normal View History

#
# 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
# ---------------------------------------------------------------------------
# Perencoder-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 QKV, 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)