Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
76 changes: 76 additions & 0 deletions tests/ut/_310p/test_model_runner_310p.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,27 @@
#
# 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.
# This file is a part of the vllm-ascend project.
#

from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock, patch

import torch
from vllm.v1.kv_cache_interface import AttentionSpec, MambaSpec

from tests.ut.base import TestBase
from vllm_ascend._310p.model_runner_310p import NPUModelRunner310


def _prepare_inputs_source() -> str:
Expand All @@ -25,3 +48,56 @@ def test_prepare_inputs_keeps_aclgraph_metadata_on_cpu() -> None:
assert "self._positions_cpu_buf[:total_num_scheduled_tokens]" in source
assert "self.seq_lens[:num_reqs].copy_(" in source
assert "self.optimistic_seq_lens_cpu[:num_reqs]" in source


class TestNPUModelRunner310(TestBase):
def test_may_reinitialize_input_batch_expands_prefix_mamba_block_table(self):
runner = object.__new__(NPUModelRunner310)
runner.max_num_reqs = 8
runner.max_model_len = 512
runner.max_encoder_len = 0
runner.max_num_tokens = 1024
runner.device = torch.device("cpu")
runner.pin_memory = False
runner.is_pooling_model = False
runner.model_config = SimpleNamespace(max_model_len=512, get_vocab_size=lambda: 32000)
runner.cache_config = SimpleNamespace(block_size=128, enable_prefix_caching=True)
runner.parallel_config = SimpleNamespace(cp_kv_cache_interleave_size=4)
runner.vllm_config = SimpleNamespace(speculative_config=None)
runner.offload_config = SimpleNamespace(uva=SimpleNamespace(cpu_offload_gb=0))
runner.input_batch = SimpleNamespace(logitsprocs=MagicMock())
attention_backend = SimpleNamespace(get_supported_kernel_block_sizes=lambda: [128, 64])
runner.attn_groups = [[SimpleNamespace(backend=attention_backend)]]

attention_spec = AttentionSpec(
block_size=128,
num_kv_heads=2,
head_size=64,
dtype=torch.float16,
)
mamba_spec = MambaSpec(
block_size=128,
shapes=((16,),),
dtypes=(torch.float16,),
mamba_cache_mode="align",
num_speculative_blocks=2,
)
kv_cache_config = SimpleNamespace(
kv_cache_groups=[
SimpleNamespace(kv_cache_spec=attention_spec),
SimpleNamespace(kv_cache_spec=mamba_spec),
]
)

with (
patch("vllm_ascend._310p.model_runner_310p.NPUInputBatch") as mock_input_batch,
patch("vllm_ascend._310p.model_runner_310p.get_total_cp_world_size", return_value=1),
):
runner.may_reinitialize_input_batch(kv_cache_config)

kwargs = mock_input_batch.call_args.kwargs
self.assertEqual(kwargs["block_sizes"], [128, 128])
self.assertEqual(kwargs["kernel_block_sizes"], [[128, 64], [0]])
self.assertEqual(kwargs["max_num_blocks_per_req"], [4, 6])
self.assertIs(kwargs["kv_cache_groups"], kv_cache_config.kv_cache_groups)
self.assertEqual(kwargs["cp_kv_cache_interleave_size"], 4)
26 changes: 12 additions & 14 deletions vllm_ascend/_310p/model_runner_310p.py
Original file line number Diff line number Diff line change
Expand Up @@ -809,24 +809,22 @@ def may_reinitialize_input_batch(self, kv_cache_config: KVCacheConfig) -> None:
Args:
kv_cache_config: The KV cache configuration.
"""
block_sizes = [
kv_cache_group.kv_cache_spec.block_size
for kv_cache_group in kv_cache_config.kv_cache_groups
if not isinstance(kv_cache_group.kv_cache_spec, EncoderOnlyAttentionSpec)
]

# Generate kernel_block_sizes that matches each block_size
# For attention backends that support virtual block splitting,
# use the supported block sizes from the backend
# For other backends (like Mamba), use [0] (no splitting)
block_sizes = []
self.kernel_block_sizes = []
kv_cache_specs = []
for kv_cache_group_id, kv_cache_group in enumerate(kv_cache_config.kv_cache_groups):
kv_cache_spec = kv_cache_group.kv_cache_spec
if isinstance(kv_cache_spec, UniformTypeKVCacheSpecs):
kv_cache_spec = next(iter(kv_cache_spec.kv_cache_specs.values()))
if isinstance(kv_cache_spec, EncoderOnlyAttentionSpec):
continue
elif isinstance(kv_cache_spec, AttentionSpec):
kv_cache_specs.append(kv_cache_spec)
block_sizes.append(kv_cache_spec.block_size)
if isinstance(kv_cache_spec, AttentionSpec):
try:
attn_groups = self.attn_groups[kv_cache_group_id]
backend = attn_groups[0].backend
Expand All @@ -845,14 +843,13 @@ def may_reinitialize_input_batch(self, kv_cache_config: KVCacheConfig) -> None:

max_num_blocks = []
max_model_len = max(self.max_model_len, self.max_encoder_len)
Comment thread
Tflowers-0129 marked this conversation as resolved.
for i, kv_cache_group in enumerate(kv_cache_config.kv_cache_groups):
if isinstance(kv_cache_group.kv_cache_spec, EncoderOnlyAttentionSpec):
continue
max_num_blocks_per_req = cdiv(max_model_len, block_sizes[i] * get_total_cp_world_size())
if isinstance(kv_cache_group.kv_cache_spec, MambaSpec):
total_cp_world_size = get_total_cp_world_size()
for kv_cache_spec in kv_cache_specs:
max_num_blocks_per_req = cdiv(max_model_len, kv_cache_spec.block_size * total_cp_world_size)
if isinstance(kv_cache_spec, MambaSpec):
mamba_blocks_per_req = (
max_num_blocks_per_req if self.cache_config.enable_prefix_caching else 1
) + kv_cache_group.kv_cache_spec.num_speculative_blocks
) + kv_cache_spec.num_speculative_blocks
max_num_blocks_per_req = max(max_num_blocks_per_req, mamba_blocks_per_req)
Comment thread
Tflowers-0129 marked this conversation as resolved.
max_num_blocks.append(max_num_blocks_per_req)

Expand All @@ -868,7 +865,7 @@ def may_reinitialize_input_batch(self, kv_cache_config: KVCacheConfig) -> None:
)
self.input_batch = NPUInputBatch(
max_num_reqs=self.max_num_reqs,
max_model_len=max(self.model_config.max_model_len, self.max_encoder_len),
max_model_len=max_model_len,
max_num_batched_tokens=self.max_num_tokens,
device=self.device,
pin_memory=self.pin_memory,
Expand All @@ -885,4 +882,5 @@ def may_reinitialize_input_batch(self, kv_cache_config: KVCacheConfig) -> None:
kernel_block_sizes=self.kernel_block_sizes,
max_num_blocks_per_req=max_num_blocks,
kv_cache_groups=kv_cache_config.kv_cache_groups,
cp_kv_cache_interleave_size=self.parallel_config.cp_kv_cache_interleave_size,
)
96 changes: 93 additions & 3 deletions vllm_ascend/patch/worker/patch_mamba_utils.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,19 @@
# mypy: ignore-errors

from typing import Any

import torch
from vllm.v1.worker import mamba_utils

from vllm_ascend.ops.triton.batch_memcpy import batch_memcpy_kernel
from vllm_ascend.utils import is_310p


def batch_memcpy(src_ptrs, dst_ptrs, sizes):
def _can_launch_triton_batch_memcpy() -> bool:
return not is_310p()


def _batch_memcpy_triton(src_ptrs, dst_ptrs, sizes):
batch = src_ptrs.shape[0]
assert dst_ptrs.shape[0] == batch
assert sizes.shape[0] == batch
Expand All @@ -17,5 +24,88 @@ def batch_memcpy(src_ptrs, dst_ptrs, sizes):
batch_memcpy_kernel[grid](src_ptrs, dst_ptrs, sizes, BLOCK_SIZE=BLOCK_SIZE)


mamba_utils.batch_memcpy_kernel = batch_memcpy_kernel
mamba_utils.batch_memcpy = batch_memcpy
def _tensor_view_from_data_ptr(state: torch.Tensor, start_addr: int, num_elements: int) -> torch.Tensor:
byte_offset = start_addr - state.data_ptr()
element_size = state.element_size()
if byte_offset < 0 or byte_offset % element_size != 0:
raise RuntimeError("Invalid Mamba state copy pointer.")

element_offset = byte_offset // element_size
flat_state = state.view(-1)
if element_offset + num_elements > flat_state.numel():
raise RuntimeError("Mamba state copy range exceeds tensor storage.")
return flat_state.narrow(0, element_offset, num_elements)


def _get_tensor_copy_pairs(copy_bufs: mamba_utils.MambaCopyBuffers) -> list[tuple[torch.Tensor, torch.Tensor]]:
if copy_bufs.offset == 0 or not hasattr(copy_bufs, "_tensor_copy_pairs"):
copy_bufs._tensor_copy_pairs = []
return copy_bufs._tensor_copy_pairs


def _collect_mamba_copy_meta_torch(
copy_bufs: mamba_utils.MambaCopyBuffers,
kv_cache_config,
mamba_state_copy_funcs,
mamba_group_ids: list[int],
src_block_idx: int,
dest_block_idx: int,
accept_token_bias: int,
req_state,
forward_context: dict[str, Any],
) -> None:
if src_block_idx == dest_block_idx and accept_token_bias == 0:
return

tensor_copy_pairs = _get_tensor_copy_pairs(copy_bufs)
sizes_np = copy_bufs.sizes.np
offset = copy_bufs.offset

for mamba_group_id in mamba_group_ids:
block_ids = req_state.block_ids[mamba_group_id]
dest_block_id = block_ids[dest_block_idx]
layer_names = kv_cache_config.kv_cache_groups[mamba_group_id].layer_names
for layer_name in layer_names:
attention = forward_context[layer_name]
kv_caches: list[torch.Tensor] = attention.kv_cache
for state, state_copy_func in zip(kv_caches, mamba_state_copy_funcs):
copy_spec = state_copy_func(state, block_ids, src_block_idx, accept_token_bias + 1)
src_state = _tensor_view_from_data_ptr(state, copy_spec.start_addr, copy_spec.num_elements)
dst_state = _tensor_view_from_data_ptr(state, state[dest_block_id].data_ptr(), copy_spec.num_elements)
tensor_copy_pairs.append((src_state, dst_state))
sizes_np[offset] = copy_spec.num_elements * state.element_size()
offset += 1

copy_bufs.offset = offset


def _do_mamba_copy_block_torch(copy_bufs: mamba_utils.MambaCopyBuffers):
n = copy_bufs.offset
if n == 0:
if hasattr(copy_bufs, "_tensor_copy_pairs"):
copy_bufs._tensor_copy_pairs = []
return

tensor_copy_pairs = getattr(copy_bufs, "_tensor_copy_pairs", None)
if tensor_copy_pairs is None or len(tensor_copy_pairs) != n:
raise RuntimeError("Mamba tensor copy metadata is incomplete.")

for src_state, dst_state in tensor_copy_pairs:
dst_state.copy_(src_state.clone())
copy_bufs._tensor_copy_pairs = []


def _batch_memcpy_unavailable(src_ptrs, dst_ptrs, sizes):
raise RuntimeError(
"Pointer-based Mamba batch memcpy requires Triton and is not available "
"on 310P. Use the tensor-copy fallback path instead."
)


if _can_launch_triton_batch_memcpy():
mamba_utils.batch_memcpy_kernel = batch_memcpy_kernel
mamba_utils.batch_memcpy = _batch_memcpy_triton
else:
mamba_utils.batch_memcpy = _batch_memcpy_unavailable
mamba_utils.collect_mamba_copy_meta = _collect_mamba_copy_meta_torch
mamba_utils.do_mamba_copy_block = _do_mamba_copy_block_torch
Loading