init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -14,5 +14,44 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from vllm_ascend.patch.platform import patch_common # noqa: F401
from vllm_ascend.patch.platform import patch_main # noqa: F401
import os
import vllm_ascend.patch.platform.patch_camem_allocator # noqa
import vllm_ascend.patch.platform.patch_distributed # noqa
import vllm_ascend.patch.platform.patch_kv_cache_utils # noqa
import vllm_ascend.patch.platform.patch_mla_prefill_backend # noqa
import vllm_ascend.patch.platform.patch_pp_mtp # noqa
import vllm_ascend.patch.platform.patch_use_v2_model_runner # noqa
from vllm_ascend.utils import is_310p, vllm_version_is
if not is_310p():
import vllm_ascend.patch.platform.patch_mamba_config # noqa
else:
import vllm_ascend.patch.platform.patch_mamba_config_310 # noqa
import vllm_ascend.patch.platform.patch_minimax_m2_config # noqa
import vllm_ascend.patch.platform.patch_glm_tool_call_streaming # noqa
if vllm_version_is("0.23.0"):
import vllm_ascend.patch.platform.patch_async_swa_kv_lifetime # noqa
import vllm_ascend.patch.platform.patch_glm47_tool_call_parser # noqa
import vllm_ascend.patch.platform.patch_minimax_m2_tool_call_parser # noqa
import vllm_ascend.patch.platform.patch_minimax_usage_accounting # noqa
import vllm_ascend.patch.platform.patch_shm_broadcast # noqa
import vllm_ascend.patch.platform.patch_deepseek_v4_tool_call_parser # noqa
import vllm_ascend.patch.platform.patch_structured_output # noqa
import vllm_ascend.patch.platform.patch_weight_transfer_engine # noqa
import vllm_ascend.patch.platform.patch_torch_accelerator # noqa
import vllm_ascend.patch.platform.patch_tool_choice_none_content # noqa
import vllm_ascend.patch.platform.patch_mamba_manager # noqa
if os.getenv("DYNAMIC_EPLB", "false").lower() in ("true", "1") or os.getenv("EXPERT_MAP_RECORD", "false") == "true":
import vllm_ascend.patch.platform.patch_multiproc_executor # noqa
import vllm_ascend.patch.platform.patch_balance_schedule # noqa
import vllm_ascend.patch.platform.patch_kv_cache_coordinator # noqa
import vllm_ascend.patch.platform.patch_speculative_config # noqa
if not vllm_version_is("0.23.0"):
import vllm_ascend.patch.platform.patch_fused_moe # noqa
import vllm_ascend.patch.platform.patch_dp_device_ids # noqa

View File

@@ -0,0 +1,190 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from contextvars import ContextVar
from functools import wraps
from typing import Any
from vllm.config import VllmConfig
from vllm.logger import logger
from vllm.v1.core.kv_cache_manager import KVCacheManager
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.core.sched.scheduler import Scheduler
from vllm.v1.core.single_type_kv_cache_manager import SingleTypeKVCacheManager
from vllm.v1.kv_cache_interface import (
ChunkedLocalAttentionSpec,
SlidingWindowSpec,
)
from vllm.v1.request import Request
_prune_context: ContextVar[tuple[str, int] | None] = ContextVar("ascend_swa_prune_context", default=None)
def _handle_negative_in_flight(
request_id: str,
num_in_flight_tokens: int,
where: str,
num_scheduled_tokens: int | None = None,
) -> None:
msg = (
"SWA_BLOCK_DIAG negative_in_flight_tokens "
f"where={where} request_id={request_id} "
f"num_in_flight_tokens={num_in_flight_tokens}"
)
if num_scheduled_tokens is not None:
msg += f" num_scheduled_tokens={num_scheduled_tokens}"
logger.warning(msg)
def _safe_in_flight_tokens(request: Request, where: str) -> int:
num_in_flight_tokens = getattr(request, "num_in_flight_tokens", 0)
if num_in_flight_tokens < 0:
_handle_negative_in_flight(
request.request_id,
num_in_flight_tokens,
where,
)
request.num_in_flight_tokens = 0
return 0
return num_in_flight_tokens
def _max_in_flight_tokens(vllm_config: VllmConfig) -> int:
return vllm_config.max_concurrent_batches * vllm_config.scheduler_config.max_num_batched_tokens
_original_request_init = Request.__init__
@wraps(_original_request_init)
def _patched_request_init(self: Request, *args: Any, **kwargs: Any) -> None:
_original_request_init(self, *args, **kwargs)
self.num_in_flight_tokens = 0
_original_update_after_schedule = Scheduler._update_after_schedule
@wraps(_original_update_after_schedule)
def _patched_update_after_schedule(self: Scheduler, scheduler_output: SchedulerOutput) -> None:
_original_update_after_schedule(self, scheduler_output)
for request_id, num_scheduled_tokens in scheduler_output.num_scheduled_tokens.items():
self.requests[request_id].num_in_flight_tokens += num_scheduled_tokens
_original_update_from_output = Scheduler.update_from_output
@wraps(_original_update_from_output)
def _patched_update_from_output(
self: Scheduler,
scheduler_output: SchedulerOutput,
model_runner_output: Any,
) -> Any:
for request_id, num_scheduled_tokens in scheduler_output.num_scheduled_tokens.items():
if request := self.requests.get(request_id):
request.num_in_flight_tokens -= num_scheduled_tokens
if request.num_in_flight_tokens < 0:
_handle_negative_in_flight(
request_id,
request.num_in_flight_tokens,
"update_from_output",
num_scheduled_tokens,
)
request.num_in_flight_tokens = 0
return _original_update_from_output(self, scheduler_output, model_runner_output)
_original_allocate_slots = KVCacheManager.allocate_slots
@wraps(_original_allocate_slots)
def _patched_allocate_slots(self: KVCacheManager, request: Request, *args: Any, **kwargs: Any) -> Any:
token = _prune_context.set((request.request_id, _safe_in_flight_tokens(request, "allocate_slots")))
try:
return _original_allocate_slots(self, request, *args, **kwargs)
finally:
_prune_context.reset(token)
_original_connector_finished = Scheduler._connector_finished
@wraps(_original_connector_finished)
def _patched_connector_finished(self: Scheduler, request: Request) -> tuple[bool, dict[str, Any] | None]:
token = _prune_context.set((request.request_id, _safe_in_flight_tokens(request, "connector_finished")))
try:
return _original_connector_finished(self, request)
finally:
_prune_context.reset(token)
_original_remove_skipped_blocks = SingleTypeKVCacheManager.remove_skipped_blocks
@wraps(_original_remove_skipped_blocks)
def _patched_remove_skipped_blocks(
self: SingleTypeKVCacheManager,
request_id: str,
total_computed_tokens: int,
) -> None:
context = _prune_context.get()
if (
context is not None
and context[0] == request_id
and isinstance(self.kv_cache_spec, (ChunkedLocalAttentionSpec, SlidingWindowSpec))
):
num_in_flight_tokens = context[1]
if num_in_flight_tokens < 0:
_handle_negative_in_flight(
request_id,
num_in_flight_tokens,
"remove_skipped_blocks",
)
num_in_flight_tokens = 0
total_computed_tokens = max(0, total_computed_tokens - num_in_flight_tokens)
_original_remove_skipped_blocks(self, request_id, total_computed_tokens)
def _patched_chunked_local_max_memory_usage_bytes(self: ChunkedLocalAttentionSpec, vllm_config: VllmConfig) -> int:
max_blocks = self.max_admission_blocks_per_request(
max_num_batched_tokens=_max_in_flight_tokens(vllm_config),
max_model_len=vllm_config.model_config.max_model_len,
)
return max_blocks * self.page_size_bytes
def _patched_swa_max_memory_usage_bytes(self: SlidingWindowSpec, vllm_config: VllmConfig) -> int:
assert vllm_config.parallel_config.decode_context_parallel_size == 1, "DCP not support sliding window."
max_blocks = self.max_admission_blocks_per_request(
max_num_batched_tokens=_max_in_flight_tokens(vllm_config),
max_model_len=vllm_config.model_config.max_model_len,
)
return max_blocks * self.page_size_bytes
_original_scheduler_init = Scheduler.__init__
@wraps(_original_scheduler_init)
def _patched_scheduler_init(self: Scheduler, vllm_config: VllmConfig, *args: Any, **kwargs: Any) -> None:
_original_scheduler_init(self, vllm_config, *args, **kwargs)
max_in_flight_tokens = _max_in_flight_tokens(vllm_config)
for manager in self.kv_cache_manager.coordinator.single_type_managers:
spec = manager.kv_cache_spec
if isinstance(spec, (ChunkedLocalAttentionSpec, SlidingWindowSpec)):
manager._max_admission_blocks_per_request = spec.max_admission_blocks_per_request(
max_num_batched_tokens=max_in_flight_tokens,
max_model_len=self.max_model_len,
)
Request.__init__ = _patched_request_init
Scheduler.__init__ = _patched_scheduler_init
Scheduler._update_after_schedule = _patched_update_after_schedule
Scheduler.update_from_output = _patched_update_from_output
Scheduler._connector_finished = _patched_connector_finished
KVCacheManager.allocate_slots = _patched_allocate_slots
SingleTypeKVCacheManager.remove_skipped_blocks = _patched_remove_skipped_blocks
ChunkedLocalAttentionSpec.max_memory_usage_bytes = _patched_chunked_local_max_memory_usage_bytes
SlidingWindowSpec.max_memory_usage_bytes = _patched_swa_max_memory_usage_bytes

View File

@@ -0,0 +1,745 @@
# mypy: ignore-errors
import os
import signal
import time
import torch
import torch.distributed as dist
import vllm
from vllm.config import ParallelConfig
from vllm.distributed.ec_transfer.ec_connector.base import ECConnectorMetadata
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorMetadata
from vllm.logger import logger
from vllm.multimodal import MULTIMODAL_REGISTRY, MultiModalRegistry
from vllm.transformers_utils.config import maybe_register_config_serialize_by_value
from vllm.utils.system_utils import decorate_logs, set_process_title
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
from vllm.v1.core.sched.interface import PauseState
from vllm.v1.core.sched.output import NewRequestData, SchedulerOutput
from vllm.v1.core.sched.request_queue import SchedulingPolicy, create_request_queue
from vllm.v1.core.sched.scheduler import Scheduler
from vllm.v1.engine import EngineCoreEventType, EngineCoreOutputs
from vllm.v1.engine.core import DPEngineCoreProc, EngineCoreProc
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.request import Request, RequestStatus
from vllm.v1.structured_output import StructuredOutputManager
from vllm.v1.utils import record_function_or_nullcontext
from vllm_ascend.utils import vllm_version_is
_ORIGINAL_RUN_ENGINE_CORE = EngineCoreProc.run_engine_core
_ORIGINAL_SCHEDULER = Scheduler
def _balance_scheduling_enabled(vllm_config) -> bool:
# TODO: Unify this path with AscendConfig once AscendConfig initialization
# is moved earlier in the startup flow.
try:
from vllm_ascend.ascend_config import get_ascend_config
return bool(get_ascend_config().enable_balance_scheduling)
except Exception:
pass
additional_config = getattr(vllm_config, "additional_config", None) or {}
if "enable_balance_scheduling" in additional_config:
return bool(additional_config["enable_balance_scheduling"])
return bool(int(os.getenv("VLLM_ASCEND_BALANCE_SCHEDULING", "0")))
def _disable_preemption_on_prefill_node(vllm_config) -> bool:
if not vllm_version_is("0.23.0"):
return False
kv_transfer_config = getattr(vllm_config, "kv_transfer_config", None)
return getattr(kv_transfer_config, "kv_role", None) == "kv_producer"
class BalanceScheduler(Scheduler):
def __init__(
self,
vllm_config,
kv_cache_config: KVCacheConfig,
structured_output_manager: StructuredOutputManager,
block_size: int,
hash_block_size: int | None = None,
mm_registry: MultiModalRegistry = MULTIMODAL_REGISTRY,
include_finished_set: bool = False,
log_stats: bool = False,
) -> None:
super().__init__(
vllm_config,
kv_cache_config,
structured_output_manager,
block_size,
hash_block_size,
mm_registry,
include_finished_set,
log_stats,
)
self._balance_enabled = _balance_scheduling_enabled(vllm_config)
self._disable_preemption = _disable_preemption_on_prefill_node(vllm_config)
if self._disable_preemption:
logger.warning("Automatic scheduler preemption is disabled on this PD-disaggregated prefill node.")
if self._balance_enabled:
self.balance_queue = [
torch.tensor([0], dtype=torch.int, device="cpu")
for _ in range(self.vllm_config.parallel_config.data_parallel_size)
]
def balance_gather(self, dp_group):
if not self._balance_enabled:
return
running_tensor = torch.tensor([len(self.running)], dtype=torch.int, device="cpu")
dist.all_gather(self.balance_queue, running_tensor, group=dp_group)
def reset_prefix_cache(
self,
reset_running_requests: bool = False,
reset_connector: bool = False,
) -> bool:
if self._disable_preemption and reset_running_requests and self.running:
raise RuntimeError(
"Cannot reset the prefix cache with running requests on a "
"PD-disaggregated prefill node because scheduler preemption "
"is disabled; drain or abort the requests first."
)
return super().reset_prefix_cache(reset_running_requests, reset_connector)
def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
if not self._balance_enabled and not self._disable_preemption:
if vllm_version_is("0.23.0"):
return super().schedule()
return super().schedule(throttle_prefills)
# NOTE(woosuk) on the scheduling algorithm:
# There's no "decoding phase" nor "prefill phase" in the scheduler.
# Each request just has the num_computed_tokens and
# num_tokens_with_spec. num_tokens_with_spec =
# len(prompt_token_ids) + len(output_token_ids) + len(spec_token_ids).
# At each step, the scheduler tries to assign tokens to the requests
# so that each request's num_computed_tokens can catch up its
# num_tokens_with_spec. This is general enough to cover
# chunked prefills, prefix caching, speculative decoding,
# and the "jump decoding" optimization in the future.
scheduled_new_reqs: list[Request] = []
scheduled_resumed_reqs: list[Request] = []
scheduled_running_reqs: list[Request] = []
preempted_reqs: list[Request] = []
req_to_new_blocks: dict[str, KVCacheBlocks] = {}
num_scheduled_tokens: dict[str, int] = {}
token_budget = self.max_num_scheduled_tokens
if self._pause_state == PauseState.PAUSED_ALL:
# Do not schedule any requests when paused.
token_budget = 0
# Encoder-related.
scheduled_encoder_inputs: dict[str, list[int]] = {}
encoder_compute_budget = self.max_num_encoder_input_tokens
# Spec decode-related.
scheduled_spec_decode_tokens: dict[str, list[int]] = {}
# For logging.
scheduled_timestamp = time.monotonic()
self.kv_cache_manager.new_step_starts()
# First, schedule the RUNNING requests.
req_index = 0
while req_index < len(self.running) and token_budget > 0:
request = self.running[req_index]
if (
request.num_output_placeholders > 0
# This is (num_computed_tokens + 1) - (num_output_placeholders - 1).
# Since output placeholders are also included in the computed tokens
# count, we subtract (num_output_placeholders - 1) to remove any draft
# tokens, so that we can be sure no further steps are needed even if
# they are all rejected.
and request.num_computed_tokens + 2 - request.num_output_placeholders
>= request.num_prompt_tokens + request.max_tokens
):
# Async scheduling: Avoid scheduling an extra step when we are sure that
# the previous step has reached request.max_tokens. We don't schedule
# partial draft tokens since this prevents uniform decode optimizations.
req_index += 1
continue
num_new_tokens = (
request.num_tokens_with_spec + request.num_output_placeholders - request.num_computed_tokens
)
if 0 < self.scheduler_config.long_prefill_token_threshold < num_new_tokens:
num_new_tokens = self.scheduler_config.long_prefill_token_threshold
num_new_tokens = min(num_new_tokens, token_budget)
# Make sure the input position does not exceed the max model len.
# This is necessary when using spec decoding.
num_new_tokens = min(num_new_tokens, self.max_model_len - 1 - request.num_computed_tokens)
# Schedule encoder inputs.
encoder_inputs_to_schedule = None
external_load_encoder_input: list[int] = []
new_encoder_compute_budget = encoder_compute_budget
if request.has_encoder_inputs:
(
encoder_inputs_to_schedule,
num_new_tokens,
new_encoder_compute_budget,
external_load_encoder_input,
) = self._try_schedule_encoder_inputs(
request,
request.num_computed_tokens,
num_new_tokens,
encoder_compute_budget,
shift_computed_tokens=1 if self.use_eagle else 0,
)
if self.need_mamba_block_aligned_split:
num_new_tokens = self._mamba_block_aligned_split(request, num_new_tokens)
if num_new_tokens == 0:
# The request cannot be scheduled because one of the following
# reasons:
# 1. No new tokens to schedule. This may happen when
# (1) PP>1 and we have already scheduled all prompt tokens
# but they are not finished yet.
# (2) Async scheduling and the request has reached to either
# its max_total_tokens or max_model_len.
# 2. The encoder budget is exhausted.
# 3. The encoder cache is exhausted.
# 4. Insufficient budget for a block-aligned chunk in hybrid
# models with mamba cache mode \"align\".
# NOTE(woosuk): Here, by doing `continue` instead of `break`,
# we do not strictly follow the FCFS scheduling policy and
# allow the lower-priority requests to be scheduled.
req_index += 1
continue
# Schedule newly needed KV blocks for the request.
with record_function_or_nullcontext("schedule: allocate_slots"):
while True:
new_blocks = self.kv_cache_manager.allocate_slots(
request,
num_new_tokens,
num_lookahead_tokens=self.num_lookahead_tokens,
)
if new_blocks is not None:
# The request can be scheduled.
break
if self._disable_preemption:
break
# The request cannot be scheduled.
# Preempt the lowest-priority request.
if self.policy == SchedulingPolicy.PRIORITY:
preempted_req = max(
self.running,
key=lambda r: (r.priority, r.arrival_time),
)
self.running.remove(preempted_req)
if preempted_req in scheduled_running_reqs:
preempted_req_id = preempted_req.request_id
scheduled_running_reqs.remove(preempted_req)
token_budget += num_scheduled_tokens.pop(preempted_req_id)
req_to_new_blocks.pop(preempted_req_id)
scheduled_spec_decode_tokens.pop(preempted_req_id, None)
preempted_encoder_inputs = scheduled_encoder_inputs.pop(preempted_req_id, None)
if preempted_encoder_inputs:
# Restore encoder compute budget if the preempted
# request had encoder inputs scheduled in this step.
num_embeds_to_restore = sum(
preempted_req.get_num_encoder_embeds(i) for i in preempted_encoder_inputs
)
encoder_compute_budget += num_embeds_to_restore
req_index -= 1
else:
preempted_req = self.running.pop()
self._preempt_request(preempted_req, scheduled_timestamp)
preempted_reqs.append(preempted_req)
if preempted_req == request:
# No more request to preempt. Cannot schedule this request.
break
if new_blocks is None:
# Cannot schedule this request.
break
# Schedule the request.
scheduled_running_reqs.append(request)
request_id = request.request_id
req_to_new_blocks[request_id] = new_blocks
num_scheduled_tokens[request_id] = num_new_tokens
token_budget -= num_new_tokens
req_index += 1
# Speculative decode related.
if request.spec_token_ids:
num_scheduled_spec_tokens = (
num_new_tokens + request.num_computed_tokens - request.num_tokens - request.num_output_placeholders
)
if num_scheduled_spec_tokens > 0:
spec_token_ids = request.spec_token_ids
if len(spec_token_ids) > num_scheduled_spec_tokens:
spec_token_ids = spec_token_ids[:num_scheduled_spec_tokens]
scheduled_spec_decode_tokens[request.request_id] = spec_token_ids
# New spec tokens will be set in `update_draft_token_ids` before the
# next step when applicable.
request.spec_token_ids = []
# Encoder-related.
if encoder_inputs_to_schedule:
scheduled_encoder_inputs[request_id] = encoder_inputs_to_schedule
# Allocate the encoder cache.
for i in encoder_inputs_to_schedule:
self.encoder_cache_manager.allocate(request, i)
encoder_compute_budget = new_encoder_compute_budget
if external_load_encoder_input:
for i in external_load_encoder_input:
self.encoder_cache_manager.allocate(request, i)
if self.ec_connector is not None:
self.ec_connector.update_state_after_alloc(request, i)
# Record the LoRAs in scheduled_running_reqs
scheduled_loras: set[int] = set()
if self.lora_config:
scheduled_loras = set(
req.lora_request.lora_int_id
for req in scheduled_running_reqs
if req.lora_request and req.lora_request.lora_int_id > 0
)
assert len(scheduled_loras) <= self.lora_config.max_loras
# Next, schedule the WAITING requests.
if not preempted_reqs and self._pause_state == PauseState.UNPAUSED:
step_skipped_waiting = create_request_queue(self.policy)
while (self.waiting or self.skipped_waiting) and token_budget > 0:
if len(self.running) == self.max_num_running_reqs:
break
if self._balance_enabled:
balance_flag = max(t.item() for t in self.balance_queue) == self.max_num_running_reqs
if balance_flag:
break
request_queue = self._select_waiting_queue_for_scheduling()
if request_queue is None:
break
request = request_queue.peek_request()
request_id = request.request_id
# try to promote blocked statuses while traversing skipped queue.
if self._is_blocked_waiting_status(request.status) and not self._try_promote_blocked_waiting_request(
request
):
if request.status == RequestStatus.WAITING_FOR_REMOTE_KVS:
logger.debug(
"%s is still in WAITING_FOR_REMOTE_KVS state.",
request_id,
)
request_queue.pop_request()
step_skipped_waiting.prepend_request(request)
continue
# Check that adding the request still respects the max_loras
# constraint.
if (
self.lora_config
and request.lora_request
and (
len(scheduled_loras) == self.lora_config.max_loras
and request.lora_request.lora_int_id not in scheduled_loras
)
):
# Scheduling would exceed max_loras, skip.
request_queue.pop_request()
step_skipped_waiting.prepend_request(request)
continue
num_external_computed_tokens = 0
load_kv_async = False
connector_prefix_cache_queries, connector_prefix_cache_hits = 0, 0
# Get already-cached tokens.
if request.num_computed_tokens == 0:
# Get locally-cached tokens.
new_computed_blocks, num_new_local_computed_tokens = self.kv_cache_manager.get_computed_blocks(
request
)
# Get externally-cached tokens if using a KVConnector.
if self.connector is not None:
ext_tokens, load_kv_async = self.connector.get_num_new_matched_tokens(
request, num_new_local_computed_tokens
)
if ext_tokens is None:
# The request cannot be scheduled because
# the KVConnector couldn't determine
# the number of matched tokens.
request_queue.pop_request()
step_skipped_waiting.prepend_request(request)
continue
num_external_computed_tokens = ext_tokens
connector_prefix_cache_queries = request.num_tokens - num_new_local_computed_tokens
connector_prefix_cache_hits = num_external_computed_tokens
# Total computed tokens (local + external).
num_computed_tokens = num_new_local_computed_tokens + num_external_computed_tokens
if request.prefill_stats is not None:
request.prefill_stats.set(
num_prompt_tokens=request.num_prompt_tokens,
num_local_cached_tokens=num_new_local_computed_tokens,
num_external_cached_tokens=num_external_computed_tokens,
)
else:
# KVTransfer: WAITING reqs have num_computed_tokens > 0
# after async KV recvs are completed.
new_computed_blocks = self.kv_cache_manager.empty_kv_cache_blocks
num_new_local_computed_tokens = 0
num_computed_tokens = request.num_computed_tokens
encoder_inputs_to_schedule = None
external_load_encoder_input = []
new_encoder_compute_budget = encoder_compute_budget
if load_kv_async:
# KVTransfer: loading remote KV, do not allocate for new work.
assert num_external_computed_tokens > 0
num_new_tokens = 0
else:
# Number of tokens to be scheduled.
# We use `request.num_tokens` instead of
# `request.num_prompt_tokens` to consider the resumed
# requests, which have output tokens.
num_new_tokens = request.num_tokens - num_computed_tokens
threshold = self.scheduler_config.long_prefill_token_threshold
if 0 < threshold < num_new_tokens:
num_new_tokens = threshold
# chunked prefill has to be enabled explicitly to allow
# pooling requests to be chunked
if not self.scheduler_config.enable_chunked_prefill and num_new_tokens > token_budget:
# If chunked_prefill is disabled,
# we can stop the scheduling here.
break
num_new_tokens = min(num_new_tokens, token_budget)
assert num_new_tokens > 0
# Schedule encoder inputs.
if request.has_encoder_inputs:
(
encoder_inputs_to_schedule,
num_new_tokens,
new_encoder_compute_budget,
external_load_encoder_input,
) = self._try_schedule_encoder_inputs(
request,
num_computed_tokens,
num_new_tokens,
encoder_compute_budget,
shift_computed_tokens=1 if self.use_eagle else 0,
)
if num_new_tokens == 0:
# The request cannot be scheduled.
break
if self.need_mamba_block_aligned_split:
num_new_tokens = self._mamba_block_aligned_split(
request,
num_new_tokens,
num_new_local_computed_tokens,
num_external_computed_tokens,
)
if num_new_tokens == 0:
break
# Handles an edge case when P/D Disaggregation
# is used with Spec Decoding where an
# extra block gets allocated which
# creates a mismatch between the number
# of local and remote blocks.
effective_lookahead_tokens = 0 if request.num_computed_tokens == 0 else self.num_lookahead_tokens
# Determine if we need to allocate cross-attention blocks.
num_encoder_tokens = 0
if self.is_encoder_decoder and request.has_encoder_inputs and encoder_inputs_to_schedule:
num_encoder_tokens = sum(request.get_num_encoder_embeds(i) for i in encoder_inputs_to_schedule)
new_blocks = self.kv_cache_manager.allocate_slots(
request,
num_new_tokens,
num_new_computed_tokens=num_new_local_computed_tokens,
new_computed_blocks=new_computed_blocks,
num_lookahead_tokens=effective_lookahead_tokens,
num_external_computed_tokens=num_external_computed_tokens,
delay_cache_blocks=load_kv_async,
num_encoder_tokens=num_encoder_tokens,
)
if new_blocks is None:
# The request cannot be scheduled.
# NOTE: we need to untouch the request from the encode cache
# manager
if request.has_encoder_inputs:
self.encoder_cache_manager.free(request)
break
# KVTransfer: the connector uses this info to determine
# if a load is needed. Note that
# This information is used to determine if a load is
# needed for this request.
if self.connector is not None:
self.connector.update_state_after_alloc(
request,
self.kv_cache_manager.get_blocks(request_id),
num_external_computed_tokens,
)
if self.connector_prefix_cache_stats is not None and connector_prefix_cache_queries != 0:
self.connector_prefix_cache_stats.record(
num_tokens=connector_prefix_cache_queries,
num_hits=connector_prefix_cache_hits,
preempted=request.num_preemptions > 0,
)
request = request_queue.pop_request()
if load_kv_async:
# If loading async, allocate memory and put request
# into the WAITING_FOR_REMOTE_KV state.
request.status = RequestStatus.WAITING_FOR_REMOTE_KVS
step_skipped_waiting.prepend_request(request)
request.num_computed_tokens = num_computed_tokens
continue
self.running.append(request)
if self.log_stats:
request.record_event(EngineCoreEventType.SCHEDULED, scheduled_timestamp)
if request.status == RequestStatus.WAITING:
scheduled_new_reqs.append(request)
elif request.status == RequestStatus.PREEMPTED:
scheduled_resumed_reqs.append(request)
else:
raise RuntimeError(f"Invalid request status: {request.status}")
if self.lora_config and request.lora_request:
scheduled_loras.add(request.lora_request.lora_int_id)
req_to_new_blocks[request_id] = self.kv_cache_manager.get_blocks(request_id)
num_scheduled_tokens[request_id] = num_new_tokens
token_budget -= num_new_tokens
request.status = RequestStatus.RUNNING
request.num_computed_tokens = num_computed_tokens
# Encoder-related.
if encoder_inputs_to_schedule:
scheduled_encoder_inputs[request_id] = encoder_inputs_to_schedule
# Allocate the encoder cache.
for i in encoder_inputs_to_schedule:
self.encoder_cache_manager.allocate(request, i)
encoder_compute_budget = new_encoder_compute_budget
# Allocate for external load encoder cache
if external_load_encoder_input:
for i in external_load_encoder_input:
self.encoder_cache_manager.allocate(request, i)
if self.ec_connector is not None:
self.ec_connector.update_state_after_alloc(request, i)
# re-queue requests skipped in this pass ahead of older skipped items.
if step_skipped_waiting:
self.skipped_waiting.prepend_requests(step_skipped_waiting)
# Check if the scheduling constraints are satisfied.
total_num_scheduled_tokens = sum(num_scheduled_tokens.values())
assert total_num_scheduled_tokens <= self.max_num_scheduled_tokens
assert token_budget >= 0
assert len(self.running) <= self.max_num_running_reqs
# Since some requests in the RUNNING queue may not be scheduled in
# this step, the total number of scheduled requests can be smaller than
# len(self.running).
assert len(scheduled_new_reqs) + len(scheduled_resumed_reqs) + len(scheduled_running_reqs) <= len(self.running)
# Get the longest common prefix among all requests in the running queue.
# This can be potentially used for cascade attention.
num_common_prefix_blocks = [0] * len(self.kv_cache_config.kv_cache_groups)
with record_function_or_nullcontext("schedule: get_num_common_prefix_blocks"):
if self.running:
any_request_id = self.running[0].request_id
num_common_prefix_blocks = self.kv_cache_manager.get_num_common_prefix_blocks(any_request_id)
# Construct the scheduler output.
if self.use_v2_model_runner:
scheduled_new_reqs = scheduled_new_reqs + scheduled_resumed_reqs
scheduled_resumed_reqs = []
new_reqs_data = [
NewRequestData.from_request(
req,
req_to_new_blocks[req.request_id].get_block_ids(),
req._all_token_ids,
)
for req in scheduled_new_reqs
]
else:
new_reqs_data = [
NewRequestData.from_request(req, req_to_new_blocks[req.request_id].get_block_ids())
for req in scheduled_new_reqs
]
with record_function_or_nullcontext("schedule: make_cached_request_data"):
cached_reqs_data = self._make_cached_request_data(
scheduled_running_reqs,
scheduled_resumed_reqs,
num_scheduled_tokens,
scheduled_spec_decode_tokens,
req_to_new_blocks,
)
# Record the request ids that were scheduled in this step.
self.prev_step_scheduled_req_ids.clear()
self.prev_step_scheduled_req_ids.update(num_scheduled_tokens.keys())
scheduler_output = SchedulerOutput(
scheduled_new_reqs=new_reqs_data,
scheduled_cached_reqs=cached_reqs_data,
num_scheduled_tokens=num_scheduled_tokens,
total_num_scheduled_tokens=total_num_scheduled_tokens,
scheduled_spec_decode_tokens=scheduled_spec_decode_tokens,
scheduled_encoder_inputs=scheduled_encoder_inputs,
num_common_prefix_blocks=num_common_prefix_blocks,
preempted_req_ids={req.request_id for req in preempted_reqs},
# finished_req_ids is an existing state in the scheduler,
# instead of being newly scheduled in this step.
# It contains the request IDs that are finished in between
# the previous and the current steps.
finished_req_ids=self.finished_req_ids,
free_encoder_mm_hashes=self.encoder_cache_manager.get_freed_mm_hashes(),
)
# NOTE(Kuntai): this function is designed for multiple purposes:
# 1. Plan the KV cache store
# 2. Wrap up all the KV cache load / save ops into an opaque object
# 3. Clear the internal states of the connector
if self.connector is not None:
meta: KVConnectorMetadata = self.connector.build_connector_meta(scheduler_output)
scheduler_output.kv_connector_metadata = meta
# Build the connector meta for ECConnector
if self.ec_connector is not None:
ec_meta: ECConnectorMetadata = self.ec_connector.build_connector_meta(scheduler_output)
scheduler_output.ec_connector_metadata = ec_meta
with record_function_or_nullcontext("schedule: update_after_schedule"):
self._update_after_schedule(scheduler_output)
return scheduler_output
class BalanceDPEngineCoreProc(DPEngineCoreProc):
def run_busy_loop(self):
"""Core busy loop of the EngineCore for data parallel case."""
# Loop until process is sent a SIGINT or SIGTERM
while True:
# 1) Poll the input queue until there is work to do.
self._process_input_queue()
# 2) Step the engine core.
executed = self._process_engine_step()
self._maybe_publish_request_counts()
local_unfinished_reqs = self.scheduler.has_unfinished_requests()
if not executed:
if not local_unfinished_reqs and not self.engines_running:
# All engines are idle.
continue
# We are in a running state and so must execute a dummy pass
# if the model didn't execute any ready requests.
self.execute_dummy_batch()
# 3) All-reduce operation to determine global unfinished reqs.
self.engines_running = self._has_global_unfinished_reqs(local_unfinished_reqs)
self.scheduler.balance_gather(self.dp_group)
if not self.engines_running:
if self.dp_rank == 0 or not self.has_coordinator:
# Notify client that we are pausing the loop.
logger.debug("Wave %d finished, pausing engine loop.", self.current_wave)
# In the coordinator case, dp rank 0 sends updates to the
# coordinator. Otherwise (offline spmd case), each rank
# sends the update to its colocated front-end process.
client_index = -1 if self.has_coordinator else 0
self.output_queue.put_nowait(
(
client_index,
EngineCoreOutputs(wave_complete=self.current_wave),
)
)
# Increment wave count and reset step counter.
self.current_wave += 1
self.step_counter = 0
def run_engine_core(*args, dp_rank: int = 0, local_dp_rank: int = 0, **kwargs):
"""Launch EngineCore busy loop in background process."""
vllm_config = kwargs.get("vllm_config")
if not _balance_scheduling_enabled(vllm_config):
return _ORIGINAL_RUN_ENGINE_CORE(*args, dp_rank=dp_rank, local_dp_rank=local_dp_rank, **kwargs)
# Signal handler used for graceful termination.
# SystemExit exception is only raised once to allow this and worker
# processes to terminate without error
shutdown_requested = False
# Ensure we can serialize transformer config after spawning
maybe_register_config_serialize_by_value()
def signal_handler(signum, frame):
nonlocal shutdown_requested
if not shutdown_requested:
shutdown_requested = True
raise SystemExit()
# Either SIGTERM or SIGINT will terminate the engine_core
signal.signal(signal.SIGTERM, signal_handler)
signal.signal(signal.SIGINT, signal_handler)
engine_core: EngineCoreProc | None = None
try:
parallel_config: ParallelConfig = kwargs["vllm_config"].parallel_config
if parallel_config.data_parallel_size > 1 or dp_rank > 0:
set_process_title("EngineCore", f"DP{dp_rank}")
decorate_logs()
# Set data parallel rank for this engine process.
parallel_config.data_parallel_rank = dp_rank
parallel_config.data_parallel_rank_local = local_dp_rank
engine_core = BalanceDPEngineCoreProc(*args, **kwargs)
else:
set_process_title("EngineCore")
decorate_logs()
engine_core = EngineCoreProc(*args, **kwargs)
engine_core.run_busy_loop()
except SystemExit:
logger.debug("EngineCore exiting.")
raise
except Exception as e:
if engine_core is None:
logger.exception("EngineCore failed to start.")
else:
logger.exception("EngineCore encountered a fatal error.")
engine_core._send_engine_dead()
raise e
finally:
if engine_core is not None:
engine_core.shutdown()
EngineCoreProc.run_engine_core = run_engine_core
vllm.v1.core.sched.scheduler.Scheduler = BalanceScheduler

View File

@@ -0,0 +1,28 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
import vllm.config.model as model_config_module
def _patched_is_cumem_allocator_available() -> bool:
# NPUPlatform declares sleep mode support and vllm-ascend uses CaMemAllocator
# in the worker path. Avoid importing the extension here because ModelConfig
# validation runs before custom op initialization.
return True
if hasattr(model_config_module, "is_cumem_allocator_available"):
model_config_module.is_cumem_allocator_available = _patched_is_cumem_allocator_available

View File

@@ -0,0 +1,802 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# DeepSeek V4 tool-call streaming parser compatibility patch.
#
from __future__ import annotations
import json
from collections import deque
from collections.abc import Sequence
from contextlib import suppress
from typing import Any
import regex as re
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
from vllm.entrypoints.openai.engine.protocol import (
DeltaFunctionCall,
DeltaMessage,
DeltaToolCall,
ExtractedToolCallInformation,
FunctionCall,
ToolCall,
)
from vllm.tool_parsers.deepseekv4_tool_parser import DeepSeekV4ToolParser
ESCAPED_ARGUMENTS_PARAM_NAME = "__vllm_param_arguments__"
def _ensure_parser_regexes(self: DeepSeekV4ToolParser) -> None:
self.tool_call_complete_regex = re.compile(
re.escape(self.tool_call_start_token) + r"(.*?)" + re.escape(self.tool_call_end_token),
re.DOTALL,
)
self.invoke_complete_regex = re.compile(
r'<|DSML|invoke\s+name="([^"]+)"\s*>(.*?)</|DSML|invoke>',
re.DOTALL,
)
self.parameter_complete_regex = re.compile(
r'<|DSML|parameter\s+name="([^"]+)"\s+string="(true|false)"\s*>(.*?)</|DSML|parameter>',
re.DOTALL,
)
self.parameter_start_regex = re.compile(r'<|DSML|parameter\s+name="([^"]+)"\s+string="(true|false)"\s*>')
self.invoke_start_regex = re.compile(r'<|DSML|invoke\s+name="([^"]+)"\s*>')
def _partial_tag_overlap(text: str, tag: str) -> int:
max_overlap = min(len(text), len(tag) - 1)
for overlap in range(max_overlap, 0, -1):
if text.endswith(tag[:overlap]):
return overlap
return 0
def _ensure_streaming_attrs(self: DeepSeekV4ToolParser) -> None:
if not hasattr(self, "_buffer"):
self._buffer = ""
if not hasattr(self, "_in_tool_calls"):
self._in_tool_calls = False
if not hasattr(self, "_active_tool_index"):
self._active_tool_index = None
if not hasattr(self, "_active_tool_name"):
self._active_tool_name = None
if not hasattr(self, "_streaming_param_mode"):
self._streaming_param_mode = None
if not hasattr(self, "_streaming_param_key"):
self._streaming_param_key = None
if not hasattr(self, "_streaming_param_raw_parts"):
self._streaming_param_raw_parts = []
if not hasattr(self, "_args_started"):
self._args_started = []
if not hasattr(self, "_pending_delta_messages"):
self._pending_delta_messages = deque()
_ensure_parser_regexes(self)
if not hasattr(self, "current_tool_index"):
self.current_tool_index = 0
if not hasattr(self, "prev_tool_call_arr"):
self.prev_tool_call_arr = []
if not hasattr(self, "streamed_args_for_tool"):
self.streamed_args_for_tool = []
def _function_name(tool) -> str | None:
if isinstance(tool, dict):
function = tool.get("function")
if isinstance(function, dict):
return function.get("name")
return getattr(function, "name", None)
return getattr(getattr(tool, "function", None), "name", None)
def _function_parameters(tool):
if isinstance(tool, dict):
function = tool.get("function")
if isinstance(function, dict):
return function.get("parameters")
return getattr(function, "parameters", None)
return getattr(getattr(tool, "function", None), "parameters", None)
def _extract_types_from_schema(schema: Any) -> list[str]:
if schema is None or not isinstance(schema, dict):
return ["string"]
types: set[str] = set()
type_value = schema.get("type")
if isinstance(type_value, str):
types.add(type_value)
elif isinstance(type_value, list):
types.update(t for t in type_value if isinstance(t, str))
enum_values = schema.get("enum")
if isinstance(enum_values, list) and enum_values:
for value in enum_values:
if value is None:
types.add("null")
elif isinstance(value, bool):
types.add("boolean")
elif isinstance(value, int):
types.add("integer")
elif isinstance(value, float):
types.add("number")
elif isinstance(value, str):
types.add("string")
elif isinstance(value, list):
types.add("array")
elif isinstance(value, dict):
types.add("object")
for choice_field in ("anyOf", "oneOf", "allOf"):
choices = schema.get(choice_field)
if isinstance(choices, list):
for choice in choices:
types.update(_extract_types_from_schema(choice))
return list(types) if types else ["string"]
_TYPE_ALIASES: dict[str, str] = {
"str": "string",
"text": "string",
"varchar": "string",
"char": "string",
"enum": "string",
"int": "integer",
"int32": "integer",
"int64": "integer",
"uint": "integer",
"uint32": "integer",
"uint64": "integer",
"long": "integer",
"short": "integer",
"unsigned": "integer",
"float": "number",
"float32": "number",
"float64": "number",
"double": "number",
"bool": "boolean",
"dict": "object",
"arr": "array",
"list": "array",
"sequence": "array",
}
def _coerce_to_schema_type(value: str, schema_type: str | list[str]) -> Any:
if isinstance(schema_type, str):
schema_type = [schema_type]
normalized_types = {_TYPE_ALIASES.get(key, key) for t in schema_type for key in [t.strip().lower()]}
for candidate_type in ("null", "integer", "number", "boolean", "object", "array", "string"):
if candidate_type not in normalized_types:
continue
if candidate_type == "null":
if value.lower() == "null":
return None
continue
if candidate_type == "string":
return value
if candidate_type == "integer":
try:
return int(value)
except (TypeError, ValueError):
continue
if candidate_type == "number":
try:
val = float(value)
return val if val != int(val) else int(val)
except (TypeError, ValueError):
continue
if candidate_type == "boolean":
lower_val = value.lower().strip()
if lower_val in ("true", "1"):
return True
if lower_val in ("false", "0"):
return False
continue
if candidate_type in ("object", "array"):
try:
return json.loads(value)
except (json.JSONDecodeError, TypeError, ValueError):
continue
try:
return json.loads(value)
except (json.JSONDecodeError, ValueError):
return value
def _convert_param_value_checked(value: str, param_type: str) -> Any:
if value.lower() == "null":
return None
param_type = param_type.lower()
if param_type in ["string", "str", "text"]:
return value
if param_type in ["integer", "int"]:
return int(value)
if param_type in ["number", "float"]:
val = float(value)
return val if val != int(val) else int(val)
if param_type in ["boolean", "bool"]:
value = value.strip()
if value.lower() not in ["false", "0", "true", "1"]:
raise ValueError("Invalid boolean value")
return value.lower() in ["true", "1"]
if param_type in ["object", "array"]:
return json.loads(value)
return json.loads(value)
def _convert_param_value(self: DeepSeekV4ToolParser, value: str, param_type) -> Any:
if not isinstance(param_type, list):
param_type = [param_type]
for current_type in param_type:
try:
return _convert_param_value_checked(value, current_type)
except Exception:
continue
return value
def _extract_param_name(param_name: str) -> str:
if param_name == ESCAPED_ARGUMENTS_PARAM_NAME:
return "arguments"
return param_name
def _get_param_config(self: DeepSeekV4ToolParser, request, function_name):
if not request or not request.tools or not function_name:
return {}
for tool in request.tools:
if _function_name(tool) != function_name:
continue
params = _function_parameters(tool)
if isinstance(params, dict):
properties = params.get("properties")
if isinstance(properties, dict):
return properties
return {}
return {}
def _coerce_param_value(
self: DeepSeekV4ToolParser,
value: str,
*,
string_attr: str,
param_type,
):
if string_attr == "true":
return value
if param_type:
return _coerce_to_schema_type(value, param_type)
try:
return json.loads(value)
except json.JSONDecodeError:
return value
def _repair_param_dict(
param_dict: dict,
param_config: dict[str, dict],
) -> dict:
allowed = set(param_config.keys())
for wrapper in ("arguments", "input"):
if set(param_dict.keys()) != {wrapper} or wrapper in allowed:
continue
inner = param_dict[wrapper]
if isinstance(inner, str):
try:
inner = json.loads(inner)
except json.JSONDecodeError:
return param_dict
if isinstance(inner, dict) and set(inner.keys()).issubset(allowed):
return inner
return param_dict
def _parse_invoke_params(
self: DeepSeekV4ToolParser,
invoke_str: str,
request: ChatCompletionRequest | None = None,
function_name: str | None = None,
) -> dict:
_ensure_parser_regexes(self)
param_config = _get_param_config(self, request, function_name)
param_dict = {}
for param_name, string_attr, param_val in self.parameter_complete_regex.findall(invoke_str):
original_param_name = param_name
param_name = _extract_param_name(param_name)
param_type = None
if original_param_name == ESCAPED_ARGUMENTS_PARAM_NAME and "arguments" in param_config:
param_type = _extract_types_from_schema(param_config["arguments"])
elif param_name in param_config and isinstance(param_config[param_name], dict):
param_type = _extract_types_from_schema(param_config[param_name])
param_dict[param_name] = _coerce_param_value(
self,
param_val,
string_attr=string_attr,
param_type=param_type,
)
return _repair_param_dict(param_dict, param_config)
def _patched_extract_tool_calls(
self: DeepSeekV4ToolParser,
model_output: str,
request: ChatCompletionRequest,
) -> ExtractedToolCallInformation:
if self.tool_call_start_token not in model_output:
return ExtractedToolCallInformation(tools_called=False, tool_calls=[], content=model_output)
try:
_ensure_parser_regexes(self)
tool_calls = []
for tool_call_match in self.tool_call_complete_regex.findall(model_output):
for invoke_name, invoke_content in self.invoke_complete_regex.findall(tool_call_match):
params = _parse_invoke_params(self, invoke_content, request, invoke_name)
tool_calls.append(
ToolCall(
type="function",
function=FunctionCall(
name=invoke_name,
arguments=json.dumps(params, ensure_ascii=False),
),
)
)
if not tool_calls:
return ExtractedToolCallInformation(tools_called=False, tool_calls=[], content=model_output)
first_tool_idx = model_output.find(self.tool_call_start_token)
content = model_output[:first_tool_idx] if first_tool_idx > 0 else None
return ExtractedToolCallInformation(
tools_called=True,
tool_calls=tool_calls,
content=content,
)
except Exception:
return ExtractedToolCallInformation(tools_called=False, tool_calls=[], content=model_output)
def _reset_streaming_state(self: DeepSeekV4ToolParser) -> None:
_ensure_streaming_attrs(self)
self.current_tool_index = 0
self._buffer = ""
self._in_tool_calls = False
self._active_tool_index = None
self._active_tool_name = None
self._streaming_param_mode = None
self._streaming_param_key = None
self._streaming_param_raw_parts.clear()
self.prev_tool_call_arr.clear()
self.streamed_args_for_tool.clear()
self._pending_delta_messages.clear()
self._args_started.clear()
def _json_escape_string_content(text: str) -> str:
return json.dumps(text, ensure_ascii=False)[1:-1]
def _drain_pending_tool_call_deltas(self: DeepSeekV4ToolParser):
while self._pending_delta_messages:
yield self._pending_delta_messages.popleft()
def _pop_pending_delta_message(self: DeepSeekV4ToolParser) -> DeltaMessage | None:
if not self._pending_delta_messages:
return None
content_parts = []
merged_tool_calls: dict[int, DeltaToolCall] = {}
while self._pending_delta_messages:
message = self._pending_delta_messages.popleft()
if message.content:
content_parts.append(message.content)
for tool_call in message.tool_calls or []:
index = tool_call.index
function = tool_call.function
if index not in merged_tool_calls:
merged_tool_calls[index] = DeltaToolCall(
index=index,
id=tool_call.id,
type=tool_call.type,
function=DeltaFunctionCall(
name=function.name if function else None,
arguments=function.arguments if function else None,
),
)
continue
merged = merged_tool_calls[index]
if tool_call.id is not None:
merged.id = tool_call.id
if tool_call.type is not None:
merged.type = tool_call.type
if function is None:
continue
if merged.function is None:
merged.function = DeltaFunctionCall()
if function.name is not None:
merged.function.name = function.name
if function.arguments is not None:
merged.function.arguments = (merged.function.arguments or "") + function.arguments
content = "".join(content_parts) or None
return DeltaMessage(content=content, tool_calls=list(merged_tool_calls.values()))
def _queue_delta_message(self: DeepSeekV4ToolParser, message: DeltaMessage | None) -> None:
if message is not None:
self._pending_delta_messages.append(message)
def _emit_tool_name_delta(self: DeepSeekV4ToolParser, index: int, name: str) -> DeltaMessage:
return DeltaMessage(
tool_calls=[
DeltaToolCall(
index=index,
id=self._generate_tool_call_id(),
function=DeltaFunctionCall(name=name, arguments=""),
type="function",
)
]
)
def _emit_tool_args_delta(self: DeepSeekV4ToolParser, index: int, arguments: str) -> DeltaMessage | None:
if not arguments:
return None
self.streamed_args_for_tool[index] += arguments
return DeltaMessage(
tool_calls=[
DeltaToolCall(
index=index,
function=DeltaFunctionCall(arguments=arguments),
)
]
)
def _begin_streaming_tool_call(self: DeepSeekV4ToolParser, name: str) -> None:
self._active_tool_index = self.current_tool_index
self._active_tool_name = name
self.current_tool_index += 1
self.prev_tool_call_arr.append({"name": name, "arguments": {}})
self.streamed_args_for_tool.append("")
self._args_started.append(False)
self._queue_delta_message(self._emit_tool_name_delta(self._active_tool_index, name))
def _append_param_prefix(self: DeepSeekV4ToolParser, index: int, key: str, *, is_string: bool) -> None:
key_json = json.dumps(key, ensure_ascii=False)
prefix = "{" if not self._args_started[index] else ","
frag = prefix + key_json + ":"
if is_string:
frag += '"'
self._args_started[index] = True
self._queue_delta_message(self._emit_tool_args_delta(index, frag))
def _append_json_param_value(self: DeepSeekV4ToolParser, index: int, key: str, value: Any) -> None:
key_json = json.dumps(key, ensure_ascii=False)
value_json = json.dumps(value, ensure_ascii=False)
prefix = "{" if not self._args_started[index] else ","
self._args_started[index] = True
self._queue_delta_message(self._emit_tool_args_delta(index, prefix + key_json + ":" + value_json))
def _append_raw_param_value(
self: DeepSeekV4ToolParser,
index: int,
key: str,
raw_value: str,
*,
is_string: bool,
) -> None:
_append_param_prefix(self, index, key, is_string=is_string)
if is_string:
frag = _json_escape_string_content(raw_value) + '"'
else:
frag = raw_value
self._queue_delta_message(self._emit_tool_args_delta(index, frag))
def _param_types_for_name(
self: DeepSeekV4ToolParser,
name: str,
request: ChatCompletionRequest | None,
) -> list[str]:
param_config = _get_param_config(self, request, self._active_tool_name)
if name in param_config and isinstance(param_config[name], dict):
return _extract_types_from_schema(param_config[name])
return ["string"]
def _can_stream_raw_param(param_types: list[str]) -> bool:
return set(param_types).issubset({"object", "array"})
def _finish_buffered_param(
self: DeepSeekV4ToolParser,
index: int,
request: ChatCompletionRequest | None,
) -> None:
key = self._streaming_param_key
if key is None:
return
raw_value = "".join(self._streaming_param_raw_parts)
param_types = _param_types_for_name(self, key, request)
value = _coerce_to_schema_type(raw_value, param_types)
_append_json_param_value(self, index, key, value)
self._streaming_param_key = None
self._streaming_param_raw_parts.clear()
def _should_buffer_wrapper_param(self: DeepSeekV4ToolParser, key: str, request: ChatCompletionRequest | None) -> bool:
if self._args_started[self._active_tool_index]:
return False
param_config = _get_param_config(self, request, self._active_tool_name)
return bool(param_config and key in ("arguments", "input") and key not in param_config)
def _finish_buffered_wrapper_param(
self: DeepSeekV4ToolParser,
index: int,
request: ChatCompletionRequest | None,
) -> None:
key = self._streaming_param_key
if key is None:
return
raw_value = "".join(self._streaming_param_raw_parts)
is_string = self._streaming_param_mode == "wrapper_string"
value: Any = raw_value
if not is_string:
try:
value = json.loads(raw_value)
except json.JSONDecodeError:
value = raw_value
param_dict = {key: value}
param_config = _get_param_config(self, request, self._active_tool_name)
repaired = _repair_param_dict(param_dict, param_config)
if isinstance(repaired, dict) and repaired is not param_dict:
for repaired_key, repaired_value in repaired.items():
_append_json_param_value(self, index, repaired_key, repaired_value)
else:
_append_raw_param_value(self, index, key, raw_value, is_string=is_string)
self._streaming_param_key = None
self._streaming_param_raw_parts.clear()
def _close_streaming_tool_call(self: DeepSeekV4ToolParser) -> None:
index = self._active_tool_index
if index is None:
return
suffix = "}" if self._args_started[index] else "{}"
self._queue_delta_message(self._emit_tool_args_delta(index, suffix))
with suppress(json.JSONDecodeError, IndexError):
self.prev_tool_call_arr[index] = {
"name": self._active_tool_name,
"arguments": json.loads(self.streamed_args_for_tool[index]),
}
self._active_tool_index = None
self._active_tool_name = None
self._streaming_param_mode = None
self._streaming_param_key = None
self._streaming_param_raw_parts.clear()
def _safe_content_len_before_tag_end(self: DeepSeekV4ToolParser) -> int:
safe_len = len(self._buffer)
parameter_end_token = "</|DSML|parameter>"
for overlap in range(1, len(parameter_end_token)):
if self._buffer.endswith(parameter_end_token[:overlap]):
safe_len = len(self._buffer) - overlap
break
return safe_len
def _process_streaming_buffer(self: DeepSeekV4ToolParser, request: ChatCompletionRequest | None) -> None:
parameter_end_token = "</|DSML|parameter>"
invoke_end_token = "</|DSML|invoke>"
while True:
if not self._in_tool_calls:
start_idx = self._buffer.find(self.tool_call_start_token)
if start_idx == -1:
overlap = _partial_tag_overlap(self._buffer, self.tool_call_start_token)
sendable_idx = len(self._buffer) - overlap
if sendable_idx > 0:
content = self._buffer[:sendable_idx]
self._buffer = self._buffer[sendable_idx:]
self._queue_delta_message(DeltaMessage(content=content))
return
if start_idx > 0:
content = self._buffer[:start_idx]
self._buffer = self._buffer[start_idx:]
self._queue_delta_message(DeltaMessage(content=content))
continue
self._buffer = self._buffer[len(self.tool_call_start_token) :]
self._in_tool_calls = True
continue
if self._active_tool_index is None:
stripped_len = len(self._buffer) - len(self._buffer.lstrip())
if stripped_len:
self._buffer = self._buffer[stripped_len:]
continue
if self._buffer.startswith(self.tool_call_end_token):
self._buffer = self._buffer[len(self.tool_call_end_token) :]
self._in_tool_calls = False
continue
match = self.invoke_start_regex.match(self._buffer)
if match is None:
return
self._buffer = self._buffer[match.end() :]
self._begin_streaming_tool_call(match.group(1))
continue
index = self._active_tool_index
if self._streaming_param_mode is not None:
end_pos = self._buffer.find(parameter_end_token)
if end_pos != -1:
raw_content = self._buffer[:end_pos]
self._buffer = self._buffer[end_pos + len(parameter_end_token) :]
if self._streaming_param_mode.startswith("wrapper_"):
self._streaming_param_raw_parts.append(raw_content)
_finish_buffered_wrapper_param(self, index, request)
elif self._streaming_param_mode == "buffered_json":
self._streaming_param_raw_parts.append(raw_content)
_finish_buffered_param(self, index, request)
elif self._streaming_param_mode == "string":
frag = _json_escape_string_content(raw_content) + '"'
self._queue_delta_message(self._emit_tool_args_delta(index, frag))
else:
frag = raw_content
self._queue_delta_message(self._emit_tool_args_delta(index, frag))
self._streaming_param_mode = None
continue
safe_len = _safe_content_len_before_tag_end(self)
if safe_len > 0:
raw_content = self._buffer[:safe_len]
self._buffer = self._buffer[safe_len:]
if self._streaming_param_mode.startswith("wrapper_") or self._streaming_param_mode == "buffered_json":
self._streaming_param_raw_parts.append(raw_content)
elif self._streaming_param_mode == "string":
frag = _json_escape_string_content(raw_content)
self._queue_delta_message(self._emit_tool_args_delta(index, frag))
else:
frag = raw_content
self._queue_delta_message(self._emit_tool_args_delta(index, frag))
return
stripped_len = len(self._buffer) - len(self._buffer.lstrip())
if stripped_len:
self._buffer = self._buffer[stripped_len:]
continue
if self._buffer.startswith(invoke_end_token):
self._buffer = self._buffer[len(invoke_end_token) :]
_close_streaming_tool_call(self)
continue
match = self.parameter_start_regex.match(self._buffer)
if match is None:
return
self._buffer = self._buffer[match.end() :]
key = _extract_param_name(match.group(1))
string_attr = match.group(2)
is_string = string_attr == "true"
if _should_buffer_wrapper_param(self, key, request):
self._streaming_param_key = key
self._streaming_param_raw_parts.clear()
self._streaming_param_mode = "wrapper_string" if is_string else "wrapper_json"
continue
if not is_string:
param_types = _param_types_for_name(self, key, request)
if not _can_stream_raw_param(param_types):
self._streaming_param_key = key
self._streaming_param_raw_parts.clear()
self._streaming_param_mode = "buffered_json"
continue
_append_param_prefix(self, index, key, is_string=is_string)
self._streaming_param_mode = "string" if is_string else "json"
def _patched_extract_tool_calls_streaming(
self: DeepSeekV4ToolParser,
previous_text: str,
current_text: str,
delta_text: str,
previous_token_ids: Sequence[int],
current_token_ids: Sequence[int],
delta_token_ids: Sequence[int],
request: ChatCompletionRequest,
) -> DeltaMessage | None:
_ensure_streaming_attrs(self)
if not previous_text:
self._reset_streaming_state()
self._buffer += delta_text
_process_streaming_buffer(self, request)
pending_delta = _pop_pending_delta_message(self)
if pending_delta is not None:
return pending_delta
if not delta_text and delta_token_ids and self.prev_tool_call_arr:
return DeltaMessage(content="")
return None
# Backward-compatible monkey patches.
DeepSeekV4ToolParser._ensure_streaming_attrs = _ensure_streaming_attrs
DeepSeekV4ToolParser._function_name = _function_name
DeepSeekV4ToolParser._function_parameters = _function_parameters
DeepSeekV4ToolParser._convert_param_value = _convert_param_value
DeepSeekV4ToolParser._extract_param_name = _extract_param_name
DeepSeekV4ToolParser._get_param_config = _get_param_config
DeepSeekV4ToolParser._coerce_param_value = _coerce_param_value
DeepSeekV4ToolParser._repair_param_dict = _repair_param_dict
DeepSeekV4ToolParser._parse_invoke_params = _parse_invoke_params
DeepSeekV4ToolParser.extract_tool_calls = _patched_extract_tool_calls
DeepSeekV4ToolParser._reset_streaming_state = _reset_streaming_state
DeepSeekV4ToolParser._json_escape_string_content = _json_escape_string_content
DeepSeekV4ToolParser.drain_pending_tool_call_deltas = _drain_pending_tool_call_deltas
DeepSeekV4ToolParser._pop_pending_delta_message = _pop_pending_delta_message
DeepSeekV4ToolParser._queue_delta_message = _queue_delta_message
DeepSeekV4ToolParser._emit_tool_name_delta = _emit_tool_name_delta
DeepSeekV4ToolParser._emit_tool_args_delta = _emit_tool_args_delta
DeepSeekV4ToolParser._begin_streaming_tool_call = _begin_streaming_tool_call
DeepSeekV4ToolParser._append_param_prefix = _append_param_prefix
DeepSeekV4ToolParser._append_json_param_value = _append_json_param_value
DeepSeekV4ToolParser._append_raw_param_value = _append_raw_param_value
DeepSeekV4ToolParser._param_types_for_name = _param_types_for_name
DeepSeekV4ToolParser._can_stream_raw_param = _can_stream_raw_param
DeepSeekV4ToolParser._finish_buffered_param = _finish_buffered_param
DeepSeekV4ToolParser._should_buffer_wrapper_param = _should_buffer_wrapper_param
DeepSeekV4ToolParser._finish_buffered_wrapper_param = _finish_buffered_wrapper_param
DeepSeekV4ToolParser._close_streaming_tool_call = _close_streaming_tool_call
DeepSeekV4ToolParser._safe_content_len_before_tag_end = _safe_content_len_before_tag_end
DeepSeekV4ToolParser._process_streaming_buffer = _process_streaming_buffer
DeepSeekV4ToolParser.extract_tool_calls_streaming = _patched_extract_tool_calls_streaming

View File

@@ -0,0 +1,89 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
#
# 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.
# Adapted from vllm/model_executor/models/qwen2_vl.py
# This file is a part of the vllm-ascend project.
import torch
from vllm_ascend.utils import AscendDeviceType, get_ascend_device_type
class NullHandle:
def __init__(self):
pass
def wait(self):
pass
def communication_adaptation_310p():
def broadcast310p_wrapper(fn):
def broadcast310p(tensor, src=0, group=None, async_op=False, group_src=None):
root = group_src if group_src is not None else src
if tensor.device == torch.device("cpu"):
return fn(tensor, src=root, group=group, async_op=async_op)
rank = torch.distributed.get_rank(group)
world_size = torch.distributed.get_world_size(group)
tensor_list = [torch.empty_like(tensor) for _ in range(world_size)]
tensor_list[rank] = tensor
torch.distributed.all_gather(tensor_list, tensor, group=group)
tensor[...] = tensor_list[src]
if async_op:
return NullHandle()
else:
return None
return broadcast310p
torch.distributed.broadcast = broadcast310p_wrapper(torch.distributed.broadcast)
torch.distributed.distributed_c10d.broadcast = broadcast310p_wrapper(torch.distributed.distributed_c10d.broadcast)
def all_reduce_wrapper_310p(fn):
def all_reduce(
tensor,
op=torch.distributed.ReduceOp.SUM,
group=None,
async_op=False,
):
if tensor.dtype != torch.int64:
return fn(tensor, op, group, async_op)
rank = torch.distributed.get_rank(group)
world_size = torch.distributed.get_world_size(group)
tensor_list = [torch.empty_like(tensor) for _ in range(world_size)]
tensor_list[rank] = tensor
torch.distributed.all_gather(tensor_list, tensor, group=group)
if op == torch.distributed.ReduceOp.SUM:
return torch.stack(tensor_list).sum(0)
elif op == torch.distributed.ReduceOp.MAX:
return torch.tensor(
torch.stack(tensor_list).cpu().numpy().max(0),
device=tensor.device,
)
else:
raise RuntimeError(f"not implement op {op}")
return all_reduce
torch.distributed.all_reduce = all_reduce_wrapper_310p(torch.distributed.all_reduce)
torch.distributed.distributed_c10d.all_reduce = all_reduce_wrapper_310p(
torch.distributed.distributed_c10d.all_reduce
)
if get_ascend_device_type() == AscendDeviceType._310P:
communication_adaptation_310p()

View File

@@ -0,0 +1,72 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# Patch vLLM v0.24.0+ ``get_physical_gpu_ids_for_local_dp_rank`` so that it
# tolerates a pre-sharded ASCEND_RT_VISIBLE_DEVICES env var (one slice per
# DP rank), instead of unconditionally applying ``local_dp_rank * world_size``
# as an offset into it.
#
# Background:
# PR #45026 removed the per-process device isolation that older vLLM
# versions performed internally. Application-level DP (e.g.
# ``offline_data_parallel.py``) now has to slice ASCEND_RT_VISIBLE_DEVICES
# per rank itself, but the upstream helper still expects the env var to
# contain ALL devices for ALL ranks and tries to read it with the
# ``local_dp_rank * world_size`` offset. With a sharded env var, that
# offset is out of range and the helper raises ``IndexError`` (wrapped in
# the user-facing "Error computing device indices for ..." message).
from vllm_ascend.utils import vllm_version_is
if not vllm_version_is("0.23.0"):
import os
from vllm.platforms import current_platform
from vllm.v1.engine import utils as _engine_utils
_original_get_physical_gpu_ids = _engine_utils.get_physical_gpu_ids_for_local_dp_rank
def _patched_get_physical_gpu_ids_for_local_dp_rank(
device_control_env_var,
local_dp_rank,
world_size,
local_world_size=None,
user_assigned_gpu_ids=None,
):
if local_world_size is None:
local_world_size = world_size
# If the caller did not pass --device-ids and the env var has
# fewer devices than the full DP range expects, the env var has
# already been pre-sharded per rank by the caller. Use it
# directly from index 0 instead of applying the DP offset again.
if user_assigned_gpu_ids is None and device_control_env_var in os.environ:
visible = [d for d in os.environ[device_control_env_var].split(",") if d]
if local_dp_rank * world_size + local_world_size > len(visible):
return [
current_platform.device_control_id_to_physical_device_id(visible[device_id])
for device_id in range(local_world_size)
]
return _original_get_physical_gpu_ids(
device_control_env_var,
local_dp_rank,
world_size,
local_world_size,
user_assigned_gpu_ids,
)
_engine_utils.get_physical_gpu_ids_for_local_dp_rank = _patched_get_physical_gpu_ids_for_local_dp_rank

View File

@@ -0,0 +1,57 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# Patch vllm's FusedMoE factory to use AscendMoERunner by default.
#
# vllm's FusedMoE is a factory function (not a class). deepseek_v2 and other
# models do `from vllm.model_executor.layers.fused_moe import FusedMoE` and
# call it directly, so we must patch the binding in the package __init__ as
# well as the layer module before any model is imported.
#
# Import order in worker.__init__:
# 1. adapt_patch() -> this file runs -> FusedMoE patched
# 2. from vllm_ascend import ops
# 3. model loading -> deepseek_v2 imported -> gets patched FusedMoE ✓
from vllm_ascend.utils import is_310p, vllm_version_is
if not vllm_version_is("0.23.0"):
import vllm.model_executor.layers.fused_moe as _fused_moe_pkg
import vllm.model_executor.layers.fused_moe.layer as _fused_moe_layer
# Capture the real original before fused_moe.py's module-level code runs.
_original_FusedMoE = _fused_moe_layer.FusedMoE
if is_310p():
from vllm_ascend._310p.fused_moe.fused_moe import AscendMoERunner310 as _DefaultAscendMoERunner
else:
from vllm_ascend.ops.fused_moe.fused_moe import AscendMoERunner as _DefaultAscendMoERunner
def _ascend_FusedMoE(*args, runner_cls=None, runner_args=None, **kwargs):
if runner_cls is None:
runner_cls = _DefaultAscendMoERunner
# 'hash' is a DeepSeek V4 flag already consumed before FusedMoE is called;
# 'tid2eid' is Ascend-specific and must reach AscendMoERunner via runner_args.
kwargs.pop("hash", None)
tid2eid = kwargs.pop("tid2eid", None)
if tid2eid is not None:
runner_args = dict(runner_args) if runner_args is not None else {}
runner_args["tid2eid"] = tid2eid
return _original_FusedMoE(*args, runner_cls=runner_cls, runner_args=runner_args, **kwargs)
_fused_moe_layer.FusedMoE = _ascend_FusedMoE
_fused_moe_pkg.FusedMoE = _ascend_FusedMoE

View File

@@ -0,0 +1,47 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# GLM-4.7 tool-call streaming parser compatibility patch.
#
from __future__ import annotations
from vllm.tool_parsers.glm47_moe_tool_parser import Glm47MoeModelToolParser
if not hasattr(Glm47MoeModelToolParser, "_ascend_original_extract_tool_call_regions"):
Glm47MoeModelToolParser._ascend_original_extract_tool_call_regions = (
Glm47MoeModelToolParser._extract_tool_call_regions
)
def _patched_extract_tool_call_regions(
self: Glm47MoeModelToolParser,
text: str,
) -> list[tuple[str, bool]]:
original_extract_tool_call_regions = self._ascend_original_extract_tool_call_regions
regions = original_extract_tool_call_regions(text)
normalized_regions: list[tuple[str, bool]] = []
for inner_text, is_complete in regions:
if is_complete and self.arg_key_start not in inner_text and "\n" not in inner_text:
tool_name = inner_text.strip()
inner_text = f"{tool_name}\n" if tool_name else inner_text
normalized_regions.append((inner_text, is_complete))
return normalized_regions
Glm47MoeModelToolParser._extract_tool_call_regions = _patched_extract_tool_call_regions

View File

@@ -0,0 +1,145 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# OpenAI chat streaming: backport GLM tool-call final chunk fixes.
#
from __future__ import annotations
import copy
import json
from typing import Any
from vllm.entrypoints.openai.chat_completion.serving import OpenAIServingChat
from vllm.entrypoints.openai.engine.protocol import (
DeltaFunctionCall,
DeltaMessage,
DeltaToolCall,
)
def _create_remaining_args_delta(
delta_message: DeltaMessage,
remaining_call: str,
index: int,
fallback_tool_call_id: str | None = None,
fallback_tool_call_type: str | None = None,
fallback_tool_call_name: str | None = None,
) -> DeltaMessage:
if remaining_call == "":
return delta_message
original_tool_call = next(
(tool_call for tool_call in delta_message.tool_calls if tool_call.index == index),
None,
)
original_function = original_tool_call.function if original_tool_call else None
function_kwargs: dict[str, str] = {"arguments": remaining_call}
function_name = original_function.name if original_function else None
if function_name is None:
function_name = fallback_tool_call_name
if function_name is not None:
function_kwargs["name"] = function_name
tool_call_kwargs: dict[str, Any] = {
"index": index,
"function": DeltaFunctionCall(**function_kwargs),
}
tool_call_id = original_tool_call.id if original_tool_call else None
if tool_call_id is None:
tool_call_id = fallback_tool_call_id
if tool_call_id is not None:
tool_call_kwargs["id"] = tool_call_id
tool_call_type = original_tool_call.type if original_tool_call else None
if tool_call_type is None:
tool_call_type = fallback_tool_call_type
if tool_call_type is not None:
tool_call_kwargs["type"] = tool_call_type
return DeltaMessage(tool_calls=[DeltaToolCall(**tool_call_kwargs)])
def _terminal_tool_arg_choice(choice: dict[str, Any]) -> bool:
if choice.get("finish_reason") != "tool_calls":
return False
delta = choice.get("delta") or {}
for tool_call in delta.get("tool_calls") or []:
function = tool_call.get("function") or {}
if function.get("arguments"):
return True
return False
def _split_terminal_tool_arg_chunk(data: str) -> list[str]:
prefix = "data: "
suffix = "\n\n"
if not data.startswith(prefix):
return [data]
payload = data[len(prefix) :]
if payload.endswith(suffix):
payload = payload[: -len(suffix)]
if payload == "[DONE]":
return [data]
try:
chunk = json.loads(payload)
except json.JSONDecodeError:
return [data]
choices = chunk.get("choices") or []
if len(choices) != 1 or not _terminal_tool_arg_choice(choices[0]):
return [data]
arg_chunk = copy.deepcopy(chunk)
arg_choice = arg_chunk["choices"][0]
arg_choice["finish_reason"] = None
arg_choice["stop_reason"] = None
finish_chunk = copy.deepcopy(chunk)
finish_choice = finish_chunk["choices"][0]
finish_choice["delta"] = {}
return [
f"{prefix}{json.dumps(arg_chunk, ensure_ascii=False)}{suffix}",
f"{prefix}{json.dumps(finish_chunk, ensure_ascii=False)}{suffix}",
]
if not hasattr(OpenAIServingChat, "_ascend_glm_original_chat_completion_stream_generator"):
OpenAIServingChat._ascend_glm_original_chat_completion_stream_generator = (
OpenAIServingChat.chat_completion_stream_generator
)
async def _wrapped_chat_completion_stream_generator(
self,
*args,
**kwargs,
):
original_stream_generator = self._ascend_glm_original_chat_completion_stream_generator
async for data in original_stream_generator(*args, **kwargs):
for chunk in _split_terminal_tool_arg_chunk(data):
yield chunk
OpenAIServingChat._create_remaining_args_delta = staticmethod(_create_remaining_args_delta)
_wrapped_chat_completion_stream_generator.__module__ = OpenAIServingChat.__module__
_wrapped_chat_completion_stream_generator.__qualname__ = (
f"{OpenAIServingChat.__qualname__}.chat_completion_stream_generator"
)
OpenAIServingChat.chat_completion_stream_generator = _wrapped_chat_completion_stream_generator

View File

@@ -0,0 +1,526 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM projectx
import sys
from collections.abc import Mapping
from math import lcm
import vllm
import vllm.envs as envs_vllm
import vllm.v1.core.kv_cache_coordinator as vllm_kv_cache_coordinator
from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_coordinator import (
HybridKVCacheCoordinator,
KVCacheCoordinator,
)
from vllm.v1.core.kv_cache_metrics import KVCacheMetricsCollector
from vllm.v1.core.kv_cache_utils import (
BlockHash,
BlockHashList,
BlockHashListWithBlockSize,
KVCacheBlock,
)
from vllm.v1.core.single_type_kv_cache_manager import (
SingleTypeKVCacheManager,
SlidingWindowManager,
)
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheConfig,
KVCacheSpec,
MambaSpec,
)
from vllm_ascend.core.single_type_kv_cache_manager import get_manager_for_kv_cache_spec
USE_MULTI_GROUPS_KV_CACHE = True
_orig_get_kv_cache_coordinator = vllm.v1.core.kv_cache_coordinator.get_kv_cache_coordinator
def _is_deepseek_v4_kv_cache_spec(kv_cache_spec: KVCacheSpec) -> bool:
if getattr(kv_cache_spec, "model_version", None) == "deepseek_v4":
return True
nested_specs = getattr(kv_cache_spec, "kv_cache_specs", None)
if nested_specs is None:
return False
if isinstance(nested_specs, Mapping):
nested_specs = nested_specs.values()
elif not isinstance(nested_specs, (list, tuple, set)):
return False
return any(getattr(spec, "model_version", None) == "deepseek_v4" for spec in nested_specs)
def _is_deepseek_v4_kv_cache_config(kv_cache_config: KVCacheConfig) -> bool:
return any(_is_deepseek_v4_kv_cache_spec(group.kv_cache_spec) for group in kv_cache_config.kv_cache_groups)
class AscendHybridKVCacheCoordinator(HybridKVCacheCoordinator):
"""
KV cache coordinator for hybrid models with multiple KV cache types, and
thus multiple kv cache groups.
To simplify `find_longest_cache_hit`, it only supports the combination of
two types of KV cache groups, and one of them must be full attention.
May extend to more general cases in the future.
"""
def __init__(
self,
kv_cache_config: KVCacheConfig,
max_model_len: int,
use_eagle: bool,
enable_caching: bool,
enable_kv_cache_events: bool,
dcp_world_size: int,
pcp_world_size: int,
hash_block_size: int,
eagle_attn_layer_names: list[str] | None = None,
metrics_collector: KVCacheMetricsCollector | None = None,
max_num_batched_tokens: int | None = None,
scheduler_block_size: int | None = None,
):
self.dcp_world_size = dcp_world_size
self.pcp_world_size = pcp_world_size
self.scheduler_block_size = scheduler_block_size
self.kv_cache_config = kv_cache_config
self.max_model_len = max_model_len
self.enable_caching = enable_caching
# Fall back to `max_model_len` when unset so the recycling-aware
# admission cap (vLLM PR #40946) collapses to the prior uncapped
# behavior. The scheduler always supplies the real value at runtime.
if max_num_batched_tokens is None:
max_num_batched_tokens = max_model_len
self.max_num_batched_tokens = max_num_batched_tokens
self.retention_interval = getattr(envs_vllm, "VLLM_PREFIX_CACHE_RETENTION_INTERVAL", None)
validate_retention_interval = getattr(
vllm_kv_cache_coordinator,
"_validate_prefix_cache_retention_interval",
None,
)
if self.retention_interval is not None and validate_retention_interval is not None:
validate_retention_interval(
self.retention_interval,
self.scheduler_block_size,
kv_cache_config,
)
self.block_pool = BlockPool(
num_gpu_blocks=kv_cache_config.num_blocks,
enable_caching=enable_caching,
hash_block_size=hash_block_size,
enable_kv_cache_events=enable_kv_cache_events,
metrics_collector=metrics_collector,
)
# KV cache group indices that get the EAGLE last-block drop.
self.eagle_group_ids: set[int] = {i for i, g in enumerate(kv_cache_config.kv_cache_groups) if g.is_eagle_group}
# Conservatively fall back to flag all groups when no group is flagged.
if use_eagle and not self.eagle_group_ids:
self.eagle_group_ids = set(range(len(kv_cache_config.kv_cache_groups)))
extra_mgr_kwargs: dict = {"scheduler_block_size": scheduler_block_size}
self.single_type_managers = tuple(
get_manager_for_kv_cache_spec(
kv_cache_spec=kv_cache_group.kv_cache_spec,
block_pool=self.block_pool,
enable_caching=enable_caching,
kv_cache_group_id=i,
dcp_world_size=dcp_world_size,
pcp_world_size=pcp_world_size,
max_num_batched_tokens=max_num_batched_tokens,
max_model_len=max_model_len,
**extra_mgr_kwargs,
)
for i, kv_cache_group in enumerate(self.kv_cache_config.kv_cache_groups)
)
# hash_block_size: the block size used to compute block hashes.
# The actual block size usually equals hash_block_size, but in cases where
# different KV cache groups have different block sizes, the actual block size
# can be a multiple of hash_block_size.
self.hash_block_size = hash_block_size
if enable_caching:
assert all(
self._get_effective_block_size(g.kv_cache_spec) % hash_block_size == 0
for g in kv_cache_config.kv_cache_groups
), "block_size must be divisible by hash_block_size"
self.verify_and_split_kv_cache_groups()
# Align the WRITE-path mask granularity (reachable_block_mask) with the
# READ-path hit granularity (find_longest_cache_hit) so SlidingWindowManager
# only caches blocks that land on a boundary where future cache hits can
# actually be matched.
# TODO (Csrayz): Consider unified all single_type_managers to simplify logic.
for mgr in self.single_type_managers:
if isinstance(mgr, SlidingWindowManager):
mgr.scheduler_block_size = self.lcm_block_size
self.use_eagle = use_eagle
def _get_effective_block_size(self, kv_cache_spec: KVCacheSpec) -> int:
block_size = kv_cache_spec.block_size
if isinstance(kv_cache_spec, MambaSpec) and self.enable_caching:
return block_size
if self.dcp_world_size * self.pcp_world_size > 1:
block_size *= self.dcp_world_size * self.pcp_world_size
if hasattr(kv_cache_spec, "compress_ratio"):
compress_ratio = kv_cache_spec.compress_ratio or 1
compress_ratio = compress_ratio if compress_ratio >= 1 else 1
block_size *= compress_ratio
return block_size
def verify_and_split_kv_cache_groups(self) -> None:
"""
Groups KV cache groups by their spec type for efficient batch processing
during cache hit lookup.
"""
attention_groups: list[tuple[KVCacheSpec, list[int], type[SingleTypeKVCacheManager]]] = []
for i, g in enumerate(self.kv_cache_config.kv_cache_groups):
manager_cls = self.single_type_managers[i].__class__
spec = g.kv_cache_spec
# Try to find an existing group with the same spec
for existing_spec, group_ids, existing_cls in attention_groups:
if existing_spec == spec:
assert manager_cls is existing_cls, "Expected same manager class for identical KV cache specs."
group_ids.append(i)
break
else:
attention_groups.append((spec, [i], manager_cls))
assert len(attention_groups) > 1, "HybridKVCacheCoordinator requires at least two attention groups."
# Put full attention first: its efficient left-to-right scan provides
# a tighter initial bound, reducing work for subsequent groups.
self.attention_groups = sorted(
attention_groups,
key=lambda x: not isinstance(x[0], FullAttentionSpec),
)
# Attention-group indices (into ``self.attention_groups``) that
# contain at least one EAGLE/MTP KV cache group.
self.eagle_attn_group_indices: set[int] = {
i
for i, (_, group_ids, _) in enumerate(self.attention_groups)
if any(gid in self.eagle_group_ids for gid in group_ids)
}
# Propagate the eagle bit to every manager in an eagle-containing
# attention group, mirroring upstream
# HybridKVCacheCoordinator.verify_and_split_kv_cache_groups. Managers
# default to ``use_eagle=False`` ("initialized lazily by the
# coordinator", see SingleTypeKVCacheManager.__init__).
#
# Required for prefix-cache correctness on DeepSeek-V4 + MTP/EAGLE: the
# SWA write path (``cache_blocks`` -> ``reachable_block_mask``) keys the
# retained checkpoint tail on ``manager.use_eagle``, while the read path
# (``find_longest_cache_hit``) applies ``drop_eagle_block`` to every gid
# merged into the eagle attention group (and ``get_cached_block``
# requires the block cached for *all* of them). If any such manager
# keeps the default False, its retained tail ends one block short of the
# eagle "peek" boundary the read looks at, the SWA group never hits, and
# the min-over-groups hybrid hit collapses to 0%. Note the upstream
# ``_annotate_eagle_groups_deepseek_v4`` flags only the single group
# holding the MTP layer, so iterating ``eagle_group_ids`` alone would
# miss its same-spec siblings.
for idx in self.eagle_attn_group_indices:
for gid in self.attention_groups[idx][1]:
self.single_type_managers[gid].use_eagle = True
# The LCM of the block sizes of all attention types.
# The cache hit length must be a multiple of the LCM of the block sizes
# to make sure the cache hit length is a multiple of the block size of
# each attention type. Requiring this because we don't support partial
# block cache hit yet.
# NOTE: use 16k as the alignment tokens for model with compress ratio
block_sizes = [self._get_effective_block_size(spec) for spec, _, _ in self.attention_groups]
self.lcm_block_size = lcm(*block_sizes)
def find_longest_cache_hit(
self,
block_hashes: list[BlockHash],
max_cache_hit_length: int,
) -> tuple[tuple[list[KVCacheBlock], ...], int]:
"""
Find the longest cache hit using an iterative fixed-point algorithm.
Each attention type either accepts the current candidate length or
reduces it. If any type reduces the length, restart checks over all
types. This converges because length monotonically decreases and is
bounded below by 0.
Args:
block_hashes: The block hashes of the request.
max_cache_hit_length: The maximum length of the cache hit.
Returns:
A tuple containing:
- A tuple of the cache hit blocks for each single type manager.
- The number of tokens of the longest cache hit.
"""
def _get_block_hashes(kv_cache_spec: KVCacheSpec) -> BlockHashList:
target_block_size = kv_cache_spec.block_size
if not isinstance(kv_cache_spec, MambaSpec) and self.dcp_world_size * self.pcp_world_size > 1:
target_block_size *= self.dcp_world_size * self.pcp_world_size
if target_block_size == self.hash_block_size:
return block_hashes
return BlockHashListWithBlockSize(block_hashes, self.hash_block_size, target_block_size)
num_groups = len(self.kv_cache_config.kv_cache_groups)
hit_length = max_cache_hit_length
hit_blocks_by_group: list[list[KVCacheBlock] | None] = [None] * num_groups
# Simple hybrid (1 full attn + 1 other): one iteration suffices.
# Full attn is always first if it exists.
is_simple_hybrid = len(self.attention_groups) == 2 and isinstance(
self.attention_groups[0][0], FullAttentionSpec
)
# Attention-group indices whose EAGLE drop is verified at the current
# ``curr_hit_length``. Each eagle group applies the drop at most once
# per candidate length (see issue #32802).
eagle_verified: set[int] = set()
while True:
curr_hit_length = hit_length
for idx, (spec, group_ids, manager_cls) in enumerate(self.attention_groups):
effective_block_size = self._get_effective_block_size(spec)
cached_blocks = hit_blocks_by_group[group_ids[0]]
if isinstance(spec, FullAttentionSpec) and cached_blocks is not None:
# Full attention is downward-closed: we only need to look
# up cached blocks once; on subsequent iterations just trim
# to the (reduced) current hit length.
num_blocks = curr_hit_length // effective_block_size
curr_hit_length = num_blocks * effective_block_size
continue
use_eagle = idx in self.eagle_attn_group_indices and idx not in eagle_verified
_max_length = curr_hit_length
if use_eagle and not isinstance(spec, MambaSpec):
# Mamba finders do not drop the EAGLE lookahead block, so
# allowing a margin here could grow the hybrid hit length.
_max_length = min(curr_hit_length + spec.block_size, max_cache_hit_length)
eagle_kwarg = {"drop_eagle_block": use_eagle}
hit_blocks = manager_cls.find_longest_cache_hit(
block_hashes=_get_block_hashes(spec),
max_length=_max_length,
kv_cache_group_ids=group_ids,
block_pool=self.block_pool,
kv_cache_spec=spec,
**eagle_kwarg,
alignment_tokens=self.lcm_block_size,
dcp_world_size=self.dcp_world_size,
pcp_world_size=self.pcp_world_size,
)
_new_hit_length = len(hit_blocks[0]) * effective_block_size
if use_eagle:
eagle_verified.add(idx)
elif _new_hit_length < curr_hit_length:
# length shrunk; invalidate previous eagle verifications
eagle_verified.clear()
curr_hit_length = _new_hit_length
curr_hit_length = len(hit_blocks[0]) * effective_block_size
for group_id, blocks in zip(group_ids, hit_blocks):
hit_blocks_by_group[group_id] = blocks
if curr_hit_length >= hit_length:
break
hit_length = curr_hit_length
if is_simple_hybrid:
break
# Truncate full attention blocks to final hit_length (if present)
# NOTE(zxr): for deepseek-v4, there is two fullattn groups, but
# in this function, only the first fullattn group is truncate by
# the belowing codes(c4), c128 layer does not truncate, which may
# have prefix cache block hit.
# Due to slidingwindow attn, deepseek-v4 decode node can't have
# any prefix cache hit, because `hit_length` of SWA is 0.
spec, group_ids, _ = self.attention_groups[0]
if isinstance(spec, FullAttentionSpec):
num_blocks = hit_length // self._get_effective_block_size(spec)
for group_id in group_ids:
if (blks := hit_blocks_by_group[group_id]) is not None:
del blks[num_blocks:]
return tuple(blocks if blocks is not None else [] for blocks in hit_blocks_by_group), hit_length
def find_longest_cache_hit_per_group(
self,
block_hashes: list[BlockHash],
max_cache_hit_length: int,
) -> tuple[tuple[list[KVCacheBlock], ...], int]:
def _get_block_hashes(kv_cache_spec: KVCacheSpec) -> BlockHashList:
target_block_size = kv_cache_spec.block_size
if not isinstance(kv_cache_spec, MambaSpec) and self.dcp_world_size * self.pcp_world_size > 1:
target_block_size *= self.dcp_world_size * self.pcp_world_size
if target_block_size == self.hash_block_size:
return block_hashes
return BlockHashListWithBlockSize(block_hashes, self.hash_block_size, target_block_size)
num_groups = len(self.kv_cache_config.kv_cache_groups)
hit_length = max_cache_hit_length
hit_blocks_by_group: list[list[KVCacheBlock] | None] = [None] * num_groups
# Simple hybrid (1 full attn + 1 other): one iteration suffices.
# Full attn is always first if it exists.
is_simple_hybrid = len(self.attention_groups) == 2 and isinstance(
self.attention_groups[0][0], FullAttentionSpec
)
# Attention-group indices whose EAGLE drop is verified at the current
# ``curr_hit_length``. Each eagle group applies the drop at most once
# per candidate length (see issue #32802).
eagle_verified: set[int] = set()
while True:
curr_hit_length = hit_length
for idx, (spec, group_ids, manager_cls) in enumerate(self.attention_groups):
# In PD disaggregation, Mamba running/temporal state is transferred
# via the KV connector, but the D side has no local Mamba prefix
# cache hit. If we let Mamba groups participate in the min-reduction,
# their zero hit collapses the FullAttention hit length to 0 and
# defeats prefix caching on the D side. Skip them instead.
if isinstance(spec, MambaSpec):
if hit_blocks_by_group[group_ids[0]] is None:
for gid in group_ids:
hit_blocks_by_group[gid] = []
continue
effective_block_size = self._get_effective_block_size(spec)
cached_blocks = hit_blocks_by_group[group_ids[0]]
if isinstance(spec, FullAttentionSpec) and cached_blocks is not None:
# Full attention is downward-closed: we only need to look
# up cached blocks once; on subsequent iterations just trim
# to the (reduced) current hit length.
num_blocks = curr_hit_length // effective_block_size
curr_hit_length = num_blocks * effective_block_size
continue
use_eagle = idx in self.eagle_attn_group_indices and idx not in eagle_verified
_max_length = curr_hit_length
if use_eagle and not isinstance(spec, MambaSpec):
# Mamba finders do not drop the EAGLE lookahead block, so
# allowing a margin here could grow the hybrid hit length.
_max_length = min(curr_hit_length + spec.block_size, max_cache_hit_length)
eagle_kwarg = {"drop_eagle_block": use_eagle}
hit_blocks = manager_cls.find_longest_cache_hit(
block_hashes=_get_block_hashes(spec),
max_length=_max_length,
kv_cache_group_ids=group_ids,
block_pool=self.block_pool,
kv_cache_spec=spec,
**eagle_kwarg,
alignment_tokens=self.lcm_block_size,
dcp_world_size=self.dcp_world_size,
pcp_world_size=self.pcp_world_size,
)
_new_hit_length = len(hit_blocks[0]) * effective_block_size
if use_eagle:
eagle_verified.add(idx)
elif _new_hit_length < curr_hit_length:
# length shrunk; invalidate previous eagle verifications
eagle_verified.clear()
curr_hit_length = _new_hit_length
curr_hit_length = len(hit_blocks[0]) * effective_block_size
for group_id, blocks in zip(group_ids, hit_blocks):
hit_blocks_by_group[group_id] = blocks
if curr_hit_length >= hit_length:
break
hit_length = curr_hit_length
if is_simple_hybrid:
break
# Truncate full attention blocks to final hit_length (if present)
# NOTE(zxr): for deepseek-v4, there is two fullattn groups, but
# in this function, only the first fullattn group is truncate by
# the belowing codes(c4), c128 layer does not truncate, which may
# have prefix cache block hit.
# Due to slidingwindow attn, deepseek-v4 decode node can't have
# any prefix cache hit, because `hit_length` of SWA is 0.
spec, group_ids, _ = self.attention_groups[0]
if isinstance(spec, FullAttentionSpec):
num_blocks = hit_length // self._get_effective_block_size(spec)
for group_id in group_ids:
if (blks := hit_blocks_by_group[group_id]) is not None:
del blks[num_blocks:]
return tuple(blocks if blocks is not None else [] for blocks in hit_blocks_by_group), hit_length
def get_kv_cache_coordinator(
kv_cache_config: KVCacheConfig,
max_model_len: int,
max_num_batched_tokens: int,
use_eagle: bool,
enable_caching: bool,
enable_kv_cache_events: bool,
dcp_world_size: int,
pcp_world_size: int,
hash_block_size: int,
scheduler_block_size: int | None = None,
eagle_attn_layer_names: list[str] | None = None,
metrics_collector: KVCacheMetricsCollector | None = None,
) -> KVCacheCoordinator:
if _is_deepseek_v4_kv_cache_config(kv_cache_config):
return AscendHybridKVCacheCoordinator(
kv_cache_config,
max_model_len,
use_eagle,
enable_caching,
enable_kv_cache_events,
dcp_world_size=dcp_world_size,
pcp_world_size=pcp_world_size,
hash_block_size=hash_block_size,
eagle_attn_layer_names=eagle_attn_layer_names,
metrics_collector=metrics_collector,
max_num_batched_tokens=max_num_batched_tokens,
scheduler_block_size=scheduler_block_size,
)
if len(kv_cache_config.kv_cache_groups) == 1 or not enable_caching:
orig_kwargs = dict(
kv_cache_config=kv_cache_config,
max_model_len=max_model_len,
max_num_batched_tokens=max_num_batched_tokens,
use_eagle=use_eagle,
enable_caching=enable_caching,
enable_kv_cache_events=enable_kv_cache_events,
dcp_world_size=dcp_world_size,
pcp_world_size=pcp_world_size,
hash_block_size=hash_block_size,
metrics_collector=metrics_collector,
)
orig_kwargs["scheduler_block_size"] = scheduler_block_size
return _orig_get_kv_cache_coordinator(**orig_kwargs)
return AscendHybridKVCacheCoordinator(
kv_cache_config,
max_model_len,
use_eagle,
enable_caching,
enable_kv_cache_events,
dcp_world_size=dcp_world_size,
pcp_world_size=pcp_world_size,
hash_block_size=hash_block_size,
eagle_attn_layer_names=eagle_attn_layer_names,
metrics_collector=metrics_collector,
max_num_batched_tokens=max_num_batched_tokens,
scheduler_block_size=scheduler_block_size,
)
vllm.v1.core.kv_cache_coordinator.get_kv_cache_coordinator = get_kv_cache_coordinator # type: ignore[attr-defined]
# `kv_cache_manager` imports `get_kv_cache_coordinator` with
# `from ... import ...`, so if it was loaded before this patch runs
# (for example through the recompute scheduler path), it keeps the
# old function object. Update that cached binding as well.
_kv_cache_manager = sys.modules.get("vllm.v1.core.kv_cache_manager")
if _kv_cache_manager is not None:
_kv_cache_manager.get_kv_cache_coordinator = get_kv_cache_coordinator # type: ignore[attr-defined]

View File

@@ -0,0 +1,359 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Ascend project
import math
from collections import defaultdict
from collections.abc import Iterable
import vllm.v1.core.block_pool
import vllm.v1.core.kv_cache_utils
from vllm.config import VllmConfig
from vllm.logger import logger
from vllm.utils.math_utils import cdiv, round_up
from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_utils import (
FreeKVCacheBlockQueue,
KVCacheBlock,
_approximate_gcd,
may_override_num_blocks,
)
from vllm.v1.kv_cache_interface import (
KVCacheConfig,
KVCacheGroupSpec,
KVCacheSpec,
KVCacheTensor,
MLAAttentionSpec,
SlidingWindowMLASpec,
UniformTypeKVCacheSpecs,
)
from vllm_ascend.utils import vllm_version_is
def _queue_block_summary(block: KVCacheBlock) -> str:
prev_id = block.prev_free_block.block_id if block.prev_free_block is not None else None
next_id = block.next_free_block.block_id if block.next_free_block is not None else None
return (
f"block_id={block.block_id} ref_cnt={block.ref_cnt} "
f"is_null={block.is_null} prev_free_block={prev_id} next_free_block={next_id}"
)
def _swa_block_diag(kind: str, block: KVCacheBlock, where: str) -> None:
msg = f"SWA_BLOCK_DIAG {kind} where={where} {_queue_block_summary(block)}"
logger.warning(msg)
def _dedupe_free_blocks(blocks: Iterable[KVCacheBlock], where: str) -> list[KVCacheBlock]:
deduped_blocks: list[KVCacheBlock] = []
seen_block_ids: set[int] = set()
for block in blocks:
if not block.is_null and block.block_id in seen_block_ids:
_swa_block_diag("duplicate_free_batch", block, where)
continue
if not block.is_null:
seen_block_ids.add(block.block_id)
deduped_blocks.append(block)
return deduped_blocks
def _filter_queue_insert_blocks(blocks: list[KVCacheBlock], where: str) -> list[KVCacheBlock]:
filtered_blocks: list[KVCacheBlock] = []
seen_block_ids: set[int] = set()
for block in blocks:
if block.is_null:
_swa_block_diag("null_free_queue_insert", block, where)
continue
if block.block_id in seen_block_ids:
_swa_block_diag("duplicate_free_queue_insert", block, where)
continue
if block.ref_cnt != 0:
_swa_block_diag("nonzero_ref_cnt_free_queue_insert", block, where)
continue
if block.prev_free_block is not None or block.next_free_block is not None:
_swa_block_diag("linked_free_queue_insert", block, where)
continue
seen_block_ids.add(block.block_id)
filtered_blocks.append(block)
return filtered_blocks
_orig_block_pool_free_blocks = BlockPool.free_blocks
def _ascend_free_blocks(
self: BlockPool,
ordered_blocks: Iterable[KVCacheBlock],
prepend: bool = False,
) -> None:
filtered_blocks: list[KVCacheBlock] = []
for block in _dedupe_free_blocks(ordered_blocks, "BlockPool.free_blocks"):
if not block.is_null and block.ref_cnt <= 0:
_swa_block_diag("ref_cnt_underflow_free_blocks", block, "BlockPool.free_blocks")
continue
filtered_blocks.append(block)
_orig_block_pool_free_blocks(self, filtered_blocks, prepend)
_orig_free_queue_prepend_n = FreeKVCacheBlockQueue.prepend_n
_orig_free_queue_append_n = FreeKVCacheBlockQueue.append_n
def _ascend_free_queue_prepend_n(self: FreeKVCacheBlockQueue, blocks: list[KVCacheBlock]) -> None:
_orig_free_queue_prepend_n(self, _filter_queue_insert_blocks(blocks, "FreeKVCacheBlockQueue.prepend_n"))
def _ascend_free_queue_append_n(self: FreeKVCacheBlockQueue, blocks: list[KVCacheBlock]) -> None:
_orig_free_queue_append_n(self, _filter_queue_insert_blocks(blocks, "FreeKVCacheBlockQueue.append_n"))
_orig_resolve_kv_cache_block_sizes = vllm.v1.core.kv_cache_utils.resolve_kv_cache_block_sizes
def _ascend_resolve_kv_cache_block_sizes(
kv_cache_config: KVCacheConfig,
vllm_config: VllmConfig,
) -> tuple[int, int]:
"""Ascend-compatible resolve_kv_cache_block_sizes.
vLLM PR #40860 added a restriction that hybrid KV cache groups with
multiple block sizes do not support context parallelism (dcp/pcp > 1).
This restriction is correct for CUDA but not for Ascend, which implements
context parallelism for MLA and SWA-MLA layers independently.
For multiple KV cache groups with CP, compute scheduler_block_size as
lcm(group_block_sizes) * dcp * pcp to maintain alignment, consistent
with the pre-PR-#40860 behavior of block_size * dcp * pcp.
"""
cache_config = vllm_config.cache_config
dcp = vllm_config.parallel_config.decode_context_parallel_size
pcp = vllm_config.parallel_config.prefill_context_parallel_size
groups = kv_cache_config.kv_cache_groups
if len(groups) <= 1:
bs = cache_config.block_size * dcp * pcp
return bs, bs
if dcp != 1 or pcp != 1:
# Ascend supports CP with multiple KV cache groups; compute
# scheduler_block_size using the LCM of all group block sizes
# multiplied by the CP factors for proper alignment.
group_block_sizes = [g.kv_cache_spec.block_size for g in groups]
scheduler_block_size = math.lcm(*group_block_sizes) * dcp * pcp
if not cache_config.enable_prefix_caching:
return scheduler_block_size, scheduler_block_size
hash_block_size = math.gcd(*group_block_sizes)
return scheduler_block_size, hash_block_size
return _orig_resolve_kv_cache_block_sizes(kv_cache_config, vllm_config)
def group_and_unify_kv_cache_specs(
kv_cache_spec: dict[str, KVCacheSpec],
) -> list[UniformTypeKVCacheSpecs] | None:
"""
Group the KV cache specs and unify each group into one UniformTypeKVCacheSpecs.
Currently, this is only used for DeepseekV4.
"""
if not any(isinstance(spec, SlidingWindowMLASpec) for spec in kv_cache_spec.values()):
return None
ratio_specs: dict[int, dict[str, KVCacheSpec]] = defaultdict(dict)
grouped_swa_mla_specs: dict[int, dict[str, KVCacheSpec]] = defaultdict(dict)
for name, spec in kv_cache_spec.items():
if isinstance(spec, SlidingWindowMLASpec):
grouped_swa_mla_specs[spec.block_size][name] = spec
elif isinstance(spec, MLAAttentionSpec):
ratio_specs[spec.compress_ratio][name] = spec
mla_uniform_specs = []
for ratio in sorted(ratio_specs, key=lambda r: (r != 4, r)):
spec_dict = ratio_specs[ratio]
assert len(spec_dict) > 0
mla_uniform_specs.append(UniformTypeKVCacheSpecs.from_specs(spec_dict))
assert mla_uniform_specs is not None
swa_uniform_specs: list[UniformTypeKVCacheSpecs] = []
for spec_dict in grouped_swa_mla_specs.values():
uniform_spec = UniformTypeKVCacheSpecs.from_specs(spec_dict)
assert uniform_spec is not None
swa_uniform_specs.append(uniform_spec)
return [*mla_uniform_specs, *swa_uniform_specs]
def _get_kv_cache_groups_uniform_groups(
grouped_specs: list[UniformTypeKVCacheSpecs],
) -> list[KVCacheGroupSpec]:
"""
Generate the KV cache groups from the grouped specs.
"""
assert len(grouped_specs) > 0 and all(isinstance(spec, UniformTypeKVCacheSpecs) for spec in grouped_specs)
# For now, we restrict the first grouped_spec to be UniformTypeKVCacheSpecs
# containing only MLAAttentionSpec.
full_mla_spec = grouped_specs[0]
full_mla_c128_spec = grouped_specs[1]
assert all(isinstance(spec, MLAAttentionSpec) for spec in full_mla_spec.kv_cache_specs.values())
full_mla_group = KVCacheGroupSpec(
layer_names=list(full_mla_spec.kv_cache_specs.keys()),
kv_cache_spec=full_mla_spec,
)
full_mla_c128_group = KVCacheGroupSpec(
layer_names=list(full_mla_c128_spec.kv_cache_specs.keys()),
kv_cache_spec=full_mla_c128_spec,
)
# We define a layer tuple as a group of layers with different page sizes, and
# one UniformTypeKVCacheSpecs contains a list of layer tuples.
# For example, if we have 11 C4 layers and 10 C128 layers, we can define a layer
# tuple as [C4I, C4A, C128], and the full_mla_group will contain "11" layer tuples.
# The other uniform KV cache specs will be similarly partitioned into layer tuples.
# Say we have 21 SWA layers, all with the same page size, then we will have "21"
# layer tuples.
num_layer_tuples_per_group: list[int] = [g_spec.get_num_layer_tuples() for g_spec in grouped_specs]
# Choose `num_layer_tuples` to minimize total padding across groups.
num_layer_tuples = _approximate_gcd(num_layer_tuples_per_group, lower_bound=num_layer_tuples_per_group[0])
# Round up to the nearest multiple of `num_layer_tuples` (i.e., padding)
num_layer_tuples_per_group = [round_up(x, num_layer_tuples) for x in num_layer_tuples_per_group]
# TODO(cmq): this is not general enough
swa_mla_specs = grouped_specs[2:]
assert all(
isinstance(spec, SlidingWindowMLASpec) for group in swa_mla_specs for spec in group.kv_cache_specs.values()
)
# Split each SWA UniformKV group into smaller groups to align their #(layer tuples)
# Possibly padding layer tuples for this.
# Additionally, we also pad KV blocks in each SWA layer, to align the page size
# with the corresponding layer in the full-MLA group.
all_page_sizes = full_mla_spec.get_page_sizes()
swa_mla_groups = []
for sm_spec in swa_mla_specs:
sm_page_sizes = sm_spec.get_page_sizes()
layers_per_size: dict[int, list[str]] = defaultdict(list)
assert max(sm_page_sizes) <= max(all_page_sizes)
# Unify page size by padding layers' page_size to the nearest larger page_size.
# Compute candidate (nearest larger page_size) for each unique page size.
size_to_candidate: dict[int, int] = {}
for ps in sm_page_sizes:
size_to_candidate[ps] = min(x for x in all_page_sizes if x >= ps)
# Pad and collect layer names per page size.
for layer_name, layer_spec in sm_spec.kv_cache_specs.items():
current_size = layer_spec.page_size_bytes
candidate = size_to_candidate[current_size]
if current_size < candidate:
object.__setattr__(layer_spec, "page_size_padded", candidate)
layers_per_size[candidate].append(layer_name)
# NOTE(yifan): for now, inside a UniformKV group, each page_size should
# have the same number of layers. This also means we don't need to pad layers
# inside a partial-full layer tuple.
assert len(set(len(layers) for layers in layers_per_size.values())) == 1
num_layers_per_size = len(next(iter(layers_per_size.values())))
# Split layers inside each UniformKV group for aligned #(layers).
# See `_get_kv_cache_groups_uniform_page_size` for more details.
num_tuple_groups = cdiv(num_layers_per_size, num_layer_tuples)
layer_tuples = list(zip(*layers_per_size.values()))
for i in range(num_tuple_groups):
group_layer_tuples = layer_tuples[i::num_tuple_groups]
# Flatten tuples and build dict for from_specs
group_layer_names = [name for layer_tuple in group_layer_tuples for name in layer_tuple]
group_layer_specs = {name: sm_spec.kv_cache_specs[name] for name in group_layer_names}
sub_sm_spec = UniformTypeKVCacheSpecs.from_specs(group_layer_specs)
assert sub_sm_spec is not None
swa_mla_groups.append(
KVCacheGroupSpec(
layer_names=group_layer_names,
kv_cache_spec=sub_sm_spec,
)
)
return [full_mla_group, full_mla_c128_group, *swa_mla_groups]
def _get_kv_cache_config_deepseek_v4(
vllm_config: VllmConfig,
kv_cache_groups: list[KVCacheGroupSpec],
available_memory: int,
) -> tuple[int, list[KVCacheTensor]]:
"""DeepseekV4 KV cache tensor layout planning.
Precondition: kv_cache_groups[0] is the full-MLA group; its page sizes
define the canonical bucket set. Non-full-MLA groups must have been
page_size-padded upstream (see _get_kv_cache_groups_uniform_groups) so
every layer's page_size matches one of the full-MLA bucket sizes.
For each group, bucket its layers by page_size_bytes and place each
layer at tuple_idx = position-within-bucket. Emit one KVCacheTensor
per (tuple_idx, bucket) whose shared_by is the union of per-group
layers at that slot.
"""
full_mla_spec = kv_cache_groups[0].kv_cache_spec
assert isinstance(full_mla_spec, UniformTypeKVCacheSpecs)
page_sizes = sorted(full_mla_spec.get_page_sizes())
layer_tuple_page_bytes = sum(page_sizes)
# Pre-bucket each group's layers by page_size (registration order within
# bucket). bucketed[g_idx][page_size] = [layer_name, ...].
mtp_layer_names = []
mtp_page_size = 0
bucketed: list[dict[int, list[str]]] = []
for group in kv_cache_groups:
assert isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs)
specs = group.kv_cache_spec.kv_cache_specs
b: dict[int, list[str]] = defaultdict(list)
for name in group.layer_names:
if "mtp" not in name:
b[specs[name].page_size_bytes].append(name)
else:
mtp_layer_names.append(name)
mtp_page_size = specs[name].page_size_bytes
bucketed.append(b)
# num_layer_tuples = longest bucket list across all groups. For the
# full-MLA group this equals the count of layers in the largest
# per-page-size bucket (= get_num_layer_tuples()); for SWA sub-groups
# this equals the sub-group size (each has a single page_size).
num_layer_tuples = max(len(layers) for b in bucketed for layers in b.values()) + len(mtp_layer_names)
num_blocks = available_memory // (layer_tuple_page_bytes * num_layer_tuples)
num_blocks = may_override_num_blocks(vllm_config, num_blocks)
kv_cache_tensors: list[KVCacheTensor] = []
for tuple_idx in range(num_layer_tuples - len(mtp_layer_names)):
for ps in page_sizes:
shared_by: list[str] = []
for b in bucketed:
bucket = b.get(ps)
if bucket is not None and tuple_idx < len(bucket):
shared_by.append(bucket[tuple_idx])
kv_cache_tensors.append(KVCacheTensor(size=ps * num_blocks, shared_by=shared_by))
for i in range(len(mtp_layer_names)):
kv_cache_tensors.append(KVCacheTensor(size=mtp_page_size * num_blocks, shared_by=[mtp_layer_names[i]]))
return num_blocks, kv_cache_tensors
BlockPool.free_blocks = _ascend_free_blocks
vllm.v1.core.block_pool.BlockPool.free_blocks = _ascend_free_blocks
FreeKVCacheBlockQueue.prepend_n = _ascend_free_queue_prepend_n
FreeKVCacheBlockQueue.append_n = _ascend_free_queue_append_n
vllm.v1.core.kv_cache_utils.FreeKVCacheBlockQueue.prepend_n = _ascend_free_queue_prepend_n
vllm.v1.core.kv_cache_utils.FreeKVCacheBlockQueue.append_n = _ascend_free_queue_append_n
vllm.v1.core.kv_cache_utils.resolve_kv_cache_block_sizes = _ascend_resolve_kv_cache_block_sizes
vllm.v1.core.kv_cache_utils.group_and_unify_kv_cache_specs = group_and_unify_kv_cache_specs
vllm.v1.core.kv_cache_utils._get_kv_cache_groups_uniform_groups = _get_kv_cache_groups_uniform_groups
# vllm v0.24.0 renamed _get_kv_cache_config_deepseek_v4 to _get_kv_cache_config_packed and
# get_kv_cache_config_from_groups now calls _get_kv_cache_config_packed directly, bypassing
# the alias patch above. Patch the canonical name so Ascend's non-packed layout is used.
if vllm_version_is("0.23.0"):
vllm.v1.core.kv_cache_utils._get_kv_cache_config_deepseek_v4 = _get_kv_cache_config_deepseek_v4
else:
vllm.v1.core.kv_cache_utils._get_kv_cache_config_packed = _get_kv_cache_config_deepseek_v4
# Also patch the reference used by engine/core.py which imports the function directly.
import vllm.v1.engine.core # noqa: E402
vllm.v1.engine.core.resolve_kv_cache_block_sizes = _ascend_resolve_kv_cache_block_sizes

View File

@@ -0,0 +1,149 @@
# mypy: ignore-errors
import math
import vllm.model_executor.models.config
from vllm.logger import logger
from vllm.model_executor.models import ModelRegistry
from vllm.model_executor.models.config import MambaModelConfig
from vllm.utils.math_utils import cdiv
from vllm.utils.torch_utils import STR_DTYPE_TO_TORCH_DTYPE, get_dtype_size
def _using_kv_store(vllm_config) -> bool:
"""
Check whether AscendStoreConnector is used.
In the scenario where only PD separation is used, mamba_cache_mode is not automatically set to align.
"""
if not vllm_config.kv_transfer_config:
return False
if vllm_config.kv_transfer_config.kv_connector == "AscendStoreConnector":
return True
if vllm_config.kv_transfer_config.kv_connector == "MultiConnector":
kv_connector_extra_config = vllm_config.kv_transfer_config.kv_connector_extra_config
if not kv_connector_extra_config:
return False
if connectors := kv_connector_extra_config.get("connectors"):
return any(connector.get("kv_connector") == "AscendStoreConnector" for connector in connectors)
return False
@classmethod
def verify_and_update_config(cls, vllm_config) -> None:
"""
Ensure that page size of attention layers is greater than or
equal to the mamba layers. If not, automatically set the attention
block size to ensure that it is. If the attention page size is
strictly greater than the mamba page size, we pad the mamba page size
to make them equal.
Args:
vllm_config: vLLM Config
"""
using_kv_store_with_hybrid = not vllm_config.scheduler_config.disable_hybrid_kv_cache_manager and _using_kv_store(
vllm_config
)
logger.debug("Using kv store: %s", using_kv_store_with_hybrid)
# Enable FULL_AND_PIECEWISE by default
MambaModelConfig.verify_and_update_config(vllm_config)
cache_config = vllm_config.cache_config
model_config = vllm_config.model_config
parallel_config = vllm_config.parallel_config
if cache_config.cache_dtype == "auto":
kv_cache_dtype = model_config.dtype
else:
kv_cache_dtype = STR_DTYPE_TO_TORCH_DTYPE[cache_config.cache_dtype]
kernel_block_size = 128
model_cls, _ = ModelRegistry.resolve_model_cls(
model_config.architecture,
model_config=model_config,
)
# get mamba block size
mamba_shapes = model_cls.get_mamba_state_shape_from_config(vllm_config)
mamba_dtypes = model_cls.get_mamba_state_dtype_from_config(vllm_config)
mamba_sizes = []
for shape, dtype in zip(mamba_shapes, mamba_dtypes):
mamba_sizes.append(math.prod(shape) * get_dtype_size(dtype))
ssm_block_page_size, conv_block_page_size = max(mamba_sizes), min(mamba_sizes)
# Pure linear attention models (e.g. bailing 2.5) have only SSM state,
# no conv block. Detected by a single 3-D mamba shape (ssm only, no conv).
# Example shape: MambaSpec(shapes=((8, 128, 128),), mamba_type='linear_attention')
if len(mamba_shapes) == 1 and len(mamba_shapes[0]) == 3:
conv_block_page_size = 0
# NOTE(zxr): because of the limit of Ascend Hardware, we need to keep
# all cache tensors contiguous, so we align the page size of ssm_block
# and single attn_block
if model_config.use_mla:
attn_num_kv_heads = model_config.get_num_kv_heads(parallel_config)
kv_lora_rank = model_config.hf_text_config.kv_lora_rank
qk_rope_head_dim = model_config.hf_text_config.qk_rope_head_dim
attn_single_token_k_page_size = kv_lora_rank * attn_num_kv_heads * get_dtype_size(kv_cache_dtype)
attn_rope_token_page_size = qk_rope_head_dim * attn_num_kv_heads * get_dtype_size(kv_cache_dtype)
attn_token_page_size = attn_single_token_k_page_size + attn_rope_token_page_size
else:
attn_num_kv_heads = model_config.get_num_kv_heads(parallel_config)
attn_head_size = model_config.get_head_size()
attn_single_token_k_page_size = attn_head_size * attn_num_kv_heads * get_dtype_size(kv_cache_dtype)
attn_token_page_size = 2 * attn_head_size * attn_num_kv_heads * get_dtype_size(kv_cache_dtype)
attn_block_size = kernel_block_size * cdiv(ssm_block_page_size, kernel_block_size * attn_single_token_k_page_size)
assert attn_single_token_k_page_size * attn_block_size == ssm_block_page_size, (
"Cannot align ssm_page_size and attn_page_size."
)
# override attention block size if either (a) the
# user has not set it or (b) the user has set it
# too small.
if cache_config.block_size is None or cache_config.block_size < attn_block_size:
cache_config.block_size = attn_block_size
logger.info(
"Setting attention block size to %d tokens to ensure that attention page size is >= mamba page size.",
attn_block_size,
)
# compute new attention page size
attn_page_size = cache_config.block_size * attn_token_page_size
# pad mamba page size for conv_blocks
if (
cache_config.mamba_page_size_padded is None
or cache_config.mamba_page_size_padded != attn_page_size + conv_block_page_size
):
cache_config.mamba_page_size_padded = attn_page_size + conv_block_page_size
mamba_padding_pct = 100 * conv_block_page_size / cache_config.mamba_page_size_padded
logger.info(
"Padding mamba page size by %.2f%% to ensure "
"that mamba page size and attention page size are "
"exactly equal.",
mamba_padding_pct,
)
# The extract_hidden_states connector (ExampleHiddenStatesConnector) only
# manages the dedicated hidden-state cache-only layer; it does not migrate
# mamba KV blocks across instances, so it does not require the block-aligned
# mamba cache mode. Forcing "align" for it would route hybrid models onto
# vLLM's fused GPU postprocess Triton kernel (introduced in vLLM #40172),
# which the Ascend Triton backend cannot compile. Leave the mode as vLLM
# derived it (e.g. "none" when prefix caching is off) for this case.
spec_config = vllm_config.speculative_config
is_extract_hidden_states = (
spec_config is not None and getattr(spec_config, "method", None) == "extract_hidden_states"
)
if using_kv_store_with_hybrid and not is_extract_hidden_states:
if cache_config.mamba_cache_mode == "none":
cache_config.mamba_cache_mode = "align"
else:
assert cache_config.mamba_cache_mode == "align", (
"mamba_cache_mode only support 'align' when kv_transfer enabled now!"
)
if cache_config.enable_prefix_caching and cache_config.mamba_cache_mode == "align":
cache_config.mamba_block_size = cache_config.block_size
else:
cache_config.mamba_block_size = model_config.max_model_len
vllm.model_executor.models.config.HybridAttentionMambaModelConfig.verify_and_update_config = verify_and_update_config

View File

@@ -0,0 +1,103 @@
# mypy: ignore-errors
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from math import lcm
import vllm.model_executor.models.config
from vllm.logger import logger
from vllm.model_executor.models import ModelRegistry
from vllm.model_executor.models.config import MambaModelConfig
from vllm.utils.math_utils import cdiv
from vllm.utils.torch_utils import STR_DTYPE_TO_TORCH_DTYPE
from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec
@classmethod
def verify_and_update_config(cls, vllm_config) -> None:
"""
Ensure that page size of attention layers is greater than or
equal to the mamba layers. If not, automatically set the attention
block size to ensure that it is. If the attention page size is
strictly greater than the mamba page size, we pad the mamba page size
to make them equal.
Args:
vllm_config: vLLM Config
"""
# Save the user input before it gets modified by MambaModelConfig
mamba_block_size = vllm_config.cache_config.mamba_block_size
# Enable FULL_AND_PIECEWISE by default
MambaModelConfig.verify_and_update_config(vllm_config)
cache_config = vllm_config.cache_config
model_config = vllm_config.model_config
parallel_config = vllm_config.parallel_config
if cache_config.cache_dtype == "auto":
kv_cache_dtype = model_config.dtype
else:
kv_cache_dtype = STR_DTYPE_TO_TORCH_DTYPE[cache_config.cache_dtype]
# get attention page size (for 1 token)
if model_config.use_mla:
raise RuntimeError("MLA is not supported on 310P currently.")
kernel_block_alignment_size = 128
attn_page_size_1_token = FullAttentionSpec(
block_size=1,
num_kv_heads=model_config.get_num_kv_heads(parallel_config),
head_size=model_config.get_head_size(),
dtype=kv_cache_dtype,
).page_size_bytes
model_cls, _ = ModelRegistry.resolve_model_cls(
model_config.architecture,
model_config=model_config,
)
# get mamba page size
mamba_page_size = MambaSpec(
shapes=model_cls.get_mamba_state_shape_from_config(vllm_config),
dtypes=model_cls.get_mamba_state_dtype_from_config(vllm_config),
block_size=-1,
).page_size_bytes
# Model may be marked as is_hybrid
# but mamba is skipped via config,
# return directly
if mamba_page_size == 0:
return
if cache_config.mamba_cache_mode == "all":
base_chunk_size = mamba_block_size or model_config.get_mamba_chunk_size()
attn_tokens_per_mamba_state = cdiv(mamba_page_size, attn_page_size_1_token)
chunk_size = lcm(base_chunk_size, kernel_block_alignment_size)
attn_block_size = chunk_size * cdiv(attn_tokens_per_mamba_state, chunk_size)
cache_config.mamba_block_size = attn_block_size
else:
attn_block_size = kernel_block_alignment_size * cdiv(
mamba_page_size, kernel_block_alignment_size * attn_page_size_1_token
)
if cache_config.block_size is None or cache_config.block_size < attn_block_size:
cache_config.block_size = attn_block_size
logger.info(
"Setting attention block size to %d tokens to ensure that attention page size is >= mamba page size.",
attn_block_size,
)
if cache_config.mamba_cache_mode == "align":
cache_config.mamba_block_size = cache_config.block_size
attn_page_size = cache_config.block_size * attn_page_size_1_token
assert attn_page_size >= mamba_page_size
if attn_page_size == mamba_page_size:
# don't need to pad mamba page size
return
# pad mamba page size to exactly match attention
if cache_config.mamba_page_size_padded is None or cache_config.mamba_page_size_padded != attn_page_size:
cache_config.mamba_page_size_padded = attn_page_size
mamba_padding_pct = 100 * (attn_page_size - mamba_page_size) / mamba_page_size
logger.info(
"Padding mamba page size by %.2f%% to ensure "
"that mamba page size and attention page size are "
"exactly equal.",
mamba_padding_pct,
)
vllm.model_executor.models.config.HybridAttentionMambaModelConfig.verify_and_update_config = verify_and_update_config

View File

@@ -0,0 +1,84 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
from collections.abc import Sequence
import vllm.v1.core.single_type_kv_cache_manager as single_type_kv_cache_manager
from vllm.v1.core.single_type_kv_cache_manager import (
BlockHashList,
BlockPool,
KVCacheBlock,
KVCacheSpec,
MambaManager,
MambaSpec,
)
class AscendMambaManager(MambaManager):
def __init__(self, kv_cache_spec: MambaSpec, block_pool: BlockPool, **kwargs) -> None:
super().__init__(kv_cache_spec, block_pool, **kwargs)
self.block_size = kv_cache_spec.block_size
@classmethod
def find_longest_cache_hit(
cls,
block_hashes: BlockHashList,
max_length: int,
kv_cache_group_ids: list[int],
block_pool: BlockPool,
kv_cache_spec: KVCacheSpec,
alignment_tokens: int,
dcp_world_size: int = 1,
pcp_world_size: int = 1,
drop_eagle_block: bool = False,
) -> tuple[list[KVCacheBlock], ...]:
assert isinstance(kv_cache_spec, MambaSpec), "MambaManager can only be used for mamba groups"
computed_blocks: tuple[list[KVCacheBlock], ...] = tuple([] for _ in range(len(kv_cache_group_ids)))
block_size = kv_cache_spec.block_size
max_num_blocks = max_length // block_size
for i in range(max_num_blocks - 1, -1, -1):
if cached_block := block_pool.get_cached_block(block_hashes[i], kv_cache_group_ids):
if block_size != alignment_tokens and (i + 1) * block_size % alignment_tokens != 0:
continue
for computed, cached in zip(computed_blocks, cached_block):
computed.extend([block_pool.null_block] * i)
computed.append(cached)
break
return computed_blocks
def get_num_blocks_to_allocate(
self,
request_id: str,
num_tokens: int,
new_computed_blocks: Sequence[KVCacheBlock],
total_computed_tokens: int,
num_tokens_main_model: int,
apply_admission_cap: bool = False,
) -> int:
num_new_blocks = super().get_num_blocks_to_allocate(
request_id,
num_tokens,
new_computed_blocks,
total_computed_tokens,
num_tokens_main_model,
apply_admission_cap,
)
# When external KV cache is loaded synchronously with new
# tokens, allocate_new_computed_blocks() allocates one
# extra block to hold the external cache content. Account
# for it here so the free-capacity check is accurate.
# (External tokens exist when total_computed_tokens exceeds
# what local prefix-cache hits cover; sync loading when
# num_tokens_main_model exceeds total_computed_tokens.)
has_external_tokens = total_computed_tokens > len(new_computed_blocks) * self.block_size
has_new_scheduled_tokens = num_tokens_main_model > total_computed_tokens
if has_external_tokens and has_new_scheduled_tokens:
# one more block for external computed tokens
num_new_blocks += 1
return num_new_blocks
single_type_kv_cache_manager.MambaManager = AscendMambaManager

View File

@@ -0,0 +1,136 @@
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# Patch target: vllm/config/model.py
# - MiniMax-M2 fp8 checkpoint on NPU: disable fp8 quantization (load bf16
# dequantized weights in worker patch) instead of failing validation.
# - For ACL graph capture, set HCCL_OP_EXPANSION_MODE=AIV if user didn't set it.
#
import os
from vllm.config.model import ModelConfig
from vllm.logger import logger
from vllm.platforms import current_platform
_original_verify_quantization = getattr(ModelConfig, "_verify_quantization", None)
_original_verify_cuda_graph = getattr(ModelConfig, "_verify_cuda_graph", None)
_DISABLE_FP8_LOG = (
"Detected fp8 MiniMax-M2 checkpoint on NPU. "
"Disabling fp8 quantization and loading dequantized bf16 "
"weights instead."
)
def _get_model_type(cfg: ModelConfig) -> str | None:
# vLLM config fields have changed across versions; try multiple sources.
model_arch_cfg = getattr(cfg, "model_arch_config", None)
if model_arch_cfg is not None:
mt = getattr(model_arch_cfg, "model_type", None)
if mt:
return mt
hf_text_cfg = getattr(cfg, "hf_text_config", None)
if hf_text_cfg is not None:
mt = getattr(hf_text_cfg, "model_type", None)
if mt:
return mt
hf_cfg = getattr(cfg, "hf_config", None)
if hf_cfg is not None:
mt = getattr(hf_cfg, "model_type", None)
if mt:
return mt
return getattr(cfg, "model_type", None)
def _should_disable_fp8(cfg: ModelConfig, quant_method: str | None) -> bool:
return current_platform.device_name == "npu" and _get_model_type(cfg) == "minimax_m2" and quant_method == "fp8"
def _disable_fp8(cfg: ModelConfig, *, log: bool) -> bool:
if not _should_disable_fp8(cfg, getattr(cfg, "quantization", None)):
return False
if log:
logger.info(_DISABLE_FP8_LOG)
cfg.quantization = None
return True
def _patched_verify_quantization(self: ModelConfig) -> None:
"""Inject mid-function behavior for ModelConfig._verify_quantization.
Upstream validates quantization inside this method via:
current_platform.verify_quantization(self.quantization)
We emulate a mid-function patch without copying upstream code by temporarily
overriding current_platform.verify_quantization while the original verifier
executes.
"""
assert _original_verify_quantization is not None
orig_platform_verify = getattr(current_platform, "verify_quantization", None)
def _platform_verify_hook(quant_method: str | None) -> None:
if _should_disable_fp8(self, quant_method):
# This is the effective "middle of _verify_quantization" interception.
_disable_fp8(self, log=True)
return
assert orig_platform_verify is not None
return orig_platform_verify(quant_method)
# Some versions may read self.quantization before calling platform verifier.
_disable_fp8(self, log=True)
try:
if orig_platform_verify is not None:
current_platform.verify_quantization = _platform_verify_hook
return _original_verify_quantization(self)
finally:
if orig_platform_verify is not None:
current_platform.verify_quantization = orig_platform_verify
# Ensure fp8 isn't restored by upstream logic.
_disable_fp8(self, log=False)
def _patched_verify_cuda_graph(self: ModelConfig) -> None:
assert _original_verify_cuda_graph is not None
if (
current_platform.device_name == "npu"
and _get_model_type(self) == "minimax_m2"
and not getattr(self, "enforce_eager", True)
):
expansion_mode = os.environ.get("HCCL_OP_EXPANSION_MODE")
if expansion_mode is None:
os.environ["HCCL_OP_EXPANSION_MODE"] = "AIV"
logger.info("Set HCCL_OP_EXPANSION_MODE=AIV for MiniMax-M2 ACL graph capture on NPU.")
elif expansion_mode != "AIV":
logger.warning(
"HCCL_OP_EXPANSION_MODE=%s may reduce ACL graph shape "
"coverage for MiniMax-M2 on NPU. Recommended value: AIV.",
expansion_mode,
)
return _original_verify_cuda_graph(self)
if _original_verify_quantization is not None:
ModelConfig._verify_quantization = _patched_verify_quantization
if _original_verify_cuda_graph is not None:
ModelConfig._verify_cuda_graph = _patched_verify_cuda_graph

View File

@@ -0,0 +1,490 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# MiniMax M2 tool parser: backport incremental tool-call argument streaming.
#
from __future__ import annotations
import json
from collections.abc import Sequence
from typing import Any
import regex as re
from vllm.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest
from vllm.entrypoints.openai.engine.protocol import (
DeltaFunctionCall,
DeltaMessage,
DeltaToolCall,
FunctionCall,
ToolCall,
)
from vllm.tokenizers import TokenizerLike
from vllm.tool_parsers import utils as tool_parser_utils
from vllm.tool_parsers.abstract_tool_parser import Tool
from vllm.tool_parsers.minimax_m2_tool_parser import MinimaxM2ToolParser
from vllm.tool_parsers.utils import (
extract_intermediate_diff,
find_tool_properties,
)
_original_init = MinimaxM2ToolParser.__init__
# vLLM main moved schema helpers from this parser class into tool_parsers.utils.
_extract_types_from_schema = getattr(tool_parser_utils, "extract_types_from_schema", None)
_coerce_to_schema_type = getattr(tool_parser_utils, "coerce_to_schema_type", None)
def _patched_init(
self: MinimaxM2ToolParser,
tokenizer: TokenizerLike,
tools: list[Tool] | None = None,
) -> None:
_original_init(self, tokenizer, tools)
tool_call_ids: list[str] = []
tool_name_sent: list[bool] = []
self._tool_call_ids = tool_call_ids
self._tool_name_sent = tool_name_sent
self._tool_call_started_from_token_id = False
def _extract_types_from_schema_fallback(schema: Any) -> list[str]:
if not isinstance(schema, dict):
return ["string"]
types: set[str] = set()
type_value = schema.get("type")
if isinstance(type_value, str):
types.add(type_value)
elif isinstance(type_value, list):
types.update(t for t in type_value if isinstance(t, str))
enum_values = schema.get("enum")
if isinstance(enum_values, list):
for value in enum_values:
if value is None:
types.add("null")
elif isinstance(value, bool):
types.add("boolean")
elif isinstance(value, int):
types.add("integer")
elif isinstance(value, float):
types.add("number")
elif isinstance(value, str):
types.add("string")
elif isinstance(value, list):
types.add("array")
elif isinstance(value, dict):
types.add("object")
for choice_field in ("anyOf", "oneOf", "allOf"):
choices = schema.get(choice_field)
if isinstance(choices, list):
for choice in choices:
types.update(_extract_types_from_schema_fallback(choice))
return list(types) if types else ["string"]
def _extract_param_types_from_schema(schema: Any) -> list[str]:
if callable(_extract_types_from_schema):
return _extract_types_from_schema(schema)
return _extract_types_from_schema_fallback(schema)
def _coerce_param_value_fallback(value: str, param_types: list[str]) -> Any:
type_aliases = {
"str": "string",
"text": "string",
"int": "integer",
"float": "number",
"bool": "boolean",
"dict": "object",
"list": "array",
}
normalized_types = {type_aliases.get(t.lower(), t.lower()) for t in param_types}
for candidate_type in ("null", "integer", "number", "boolean", "object", "array", "string"):
if candidate_type not in normalized_types:
continue
if candidate_type == "null":
if value.lower() == "null":
return None
continue
if candidate_type == "string":
return value
if candidate_type == "integer":
try:
return int(value)
except (ValueError, TypeError):
continue
if candidate_type == "number":
try:
val = float(value)
return val if val != int(val) else int(val)
except (ValueError, TypeError):
continue
if candidate_type == "boolean":
lower_val = value.lower().strip()
if lower_val in ("true", "1"):
return True
if lower_val in ("false", "0"):
return False
continue
if candidate_type in ("object", "array"):
try:
return json.loads(value)
except (json.JSONDecodeError, ValueError, TypeError):
continue
try:
return json.loads(value)
except (json.JSONDecodeError, ValueError):
return value
def _coerce_param_value(value: str, param_types: list[str]) -> Any:
if callable(_coerce_to_schema_type):
return _coerce_to_schema_type(value, param_types)
return _coerce_param_value_fallback(value, param_types)
def _get_param_types_from_config(
param_name: str,
param_config: dict[str, Any],
) -> list[str]:
param_schema = param_config.get(param_name)
if not isinstance(param_schema, dict):
return ["string"]
return _extract_param_types_from_schema(param_schema)
def _patched_parse_single_invoke(
self: MinimaxM2ToolParser,
invoke_str: str,
tools: list[Tool] | None,
) -> ToolCall | None:
name_match = re.search(r"^([^>]+)", invoke_str)
if not name_match:
return None
function_name = self._extract_name(name_match.group(1))
param_config = find_tool_properties(tools, function_name)
param_dict = {}
for match in self.parameter_complete_regex.findall(invoke_str):
param_match = re.search(r"^([^>]+)>(.*)", match, re.DOTALL)
if param_match:
param_name = self._extract_name(param_match.group(1))
param_value = param_match.group(2).strip()
param_type = _get_param_types_from_config(param_name, param_config)
param_dict[param_name] = _coerce_param_value(param_value, param_type)
return ToolCall(
type="function",
function=FunctionCall(
name=function_name,
arguments=json.dumps(param_dict, ensure_ascii=False),
),
)
def _reset_streaming_state(
self: MinimaxM2ToolParser,
tool_call_started: bool = False,
) -> None:
self.current_tool_index = 0
self.prev_tool_call_arr.clear()
self.streamed_args_for_tool.clear()
self._tool_call_ids.clear()
self._tool_name_sent.clear()
self._tool_call_started_from_token_id = False
self.is_tool_call_started = tool_call_started
def _ensure_streaming_slots(self: MinimaxM2ToolParser, tool_count: int) -> None:
while len(self.streamed_args_for_tool) < tool_count:
self.streamed_args_for_tool.append("")
while len(self._tool_call_ids) < tool_count:
self._tool_call_ids.append(self._generate_tool_call_id())
while len(self._tool_name_sent) < tool_count:
self._tool_name_sent.append(False)
def _get_param_config(
self: MinimaxM2ToolParser,
function_name: str,
) -> dict[str, Any]:
return find_tool_properties(self.tools, function_name)
def _serialize_partial_param_value(
self: MinimaxM2ToolParser,
value: str,
param_types: list[str],
) -> str:
value = value.strip()
converted = _coerce_param_value(value, param_types)
return json.dumps(converted, ensure_ascii=False)
def _build_partial_arguments(
self: MinimaxM2ToolParser,
invoke_body: str,
*,
invoke_complete: bool,
param_config: dict[str, Any],
) -> str:
args_parts: list[str] = []
search_pos = 0
while True:
param_start = invoke_body.find("<parameter name=", search_pos)
if param_start == -1:
break
name_start = param_start + len("<parameter name=")
name_end = invoke_body.find(">", name_start)
if name_end == -1:
break
param_name = self._extract_name(invoke_body[name_start:name_end])
value_start = name_end + 1
value_end = invoke_body.find("</parameter>", value_start)
param_complete = value_end != -1
if not param_complete:
break
param_value = invoke_body[value_start:value_end]
search_pos = value_end + len("</parameter>")
param_types = _get_param_types_from_config(param_name, param_config)
serialized_value = self._serialize_partial_param_value(
param_value,
param_types,
)
if not serialized_value:
break
args_parts.append(f"{json.dumps(param_name, ensure_ascii=False)}:{serialized_value}")
if not args_parts:
return "{}" if invoke_complete else ""
args_json = "{" + ",".join(args_parts)
if invoke_complete:
args_json += "}"
return args_json
def _get_invoke_states(
self: MinimaxM2ToolParser,
current_text: str,
) -> list[dict[str, Any]]:
tool_start = current_text.find(self.tool_call_start_token)
if tool_start == -1:
if not self.is_tool_call_started:
return []
tool_payload = current_text
else:
tool_payload = current_text[tool_start + len(self.tool_call_start_token) :]
tool_end = tool_payload.find(self.tool_call_end_token)
if tool_end != -1:
tool_payload = tool_payload[:tool_end]
invoke_states: list[dict[str, Any]] = []
search_pos = 0
while True:
invoke_start = tool_payload.find("<invoke name=", search_pos)
if invoke_start == -1:
break
invoke_content_start = invoke_start + len("<invoke name=")
invoke_end = tool_payload.find("</invoke>", invoke_content_start)
invoke_complete = invoke_end != -1
if invoke_complete:
invoke_str = tool_payload[invoke_content_start:invoke_end]
search_pos = invoke_end + len("</invoke>")
else:
invoke_str = tool_payload[invoke_content_start:]
search_pos = len(tool_payload)
name_end = invoke_str.find(">")
if name_end == -1:
break
function_name = self._extract_name(invoke_str[:name_end])
param_config = self._get_param_config(function_name)
invoke_body = invoke_str[name_end + 1 :]
partial_args = self._build_partial_arguments(
invoke_body,
invoke_complete=invoke_complete,
param_config=param_config,
)
tool_call = self._parse_single_invoke(invoke_str, self.tools) if invoke_complete else None
invoke_states.append(
{
"name": function_name,
"arguments": partial_args,
"complete": invoke_complete,
"tool_call": tool_call,
}
)
if not invoke_complete:
break
return invoke_states
def _finalize_completed_tool_call(
self: MinimaxM2ToolParser,
idx: int,
invoke_state: dict[str, Any],
) -> None:
if not invoke_state["complete"] or len(self.prev_tool_call_arr) > idx:
return
tool_call = invoke_state["tool_call"]
if tool_call is None:
return
self.prev_tool_call_arr.append(
{
"name": tool_call.function.name,
"arguments": json.loads(tool_call.function.arguments),
}
)
def _extract_delta_tool_call(
self: MinimaxM2ToolParser,
current_text: str,
) -> DeltaToolCall | None:
invoke_states = self._get_invoke_states(current_text)
if not invoke_states:
return None
self._ensure_streaming_slots(len(invoke_states))
for idx, invoke_state in enumerate(invoke_states):
args_json = invoke_state["arguments"]
sent_args = self.streamed_args_for_tool[idx]
name_sent = self._tool_name_sent[idx]
if not name_sent:
self._tool_name_sent[idx] = True
self.current_tool_index = idx
if args_json:
self.streamed_args_for_tool[idx] = args_json
self._finalize_completed_tool_call(idx, invoke_state)
return DeltaToolCall(
index=idx,
id=self._tool_call_ids[idx],
type="function",
function=DeltaFunctionCall(
name=invoke_state["name"],
arguments=args_json or None,
),
)
if args_json and args_json != sent_args:
if sent_args and args_json.startswith(sent_args):
args_delta = args_json[len(sent_args) :]
else:
args_delta = extract_intermediate_diff(args_json, sent_args)
if args_delta:
self.streamed_args_for_tool[idx] = args_json
self.current_tool_index = idx
self._finalize_completed_tool_call(idx, invoke_state)
return DeltaToolCall(
index=idx,
function=DeltaFunctionCall(arguments=args_delta),
)
self._finalize_completed_tool_call(idx, invoke_state)
return None
def _patched_extract_tool_calls_streaming(
self: MinimaxM2ToolParser,
previous_text: str,
current_text: str,
delta_text: str,
previous_token_ids: Sequence[int], # pylint: disable=unused-argument
current_token_ids: Sequence[int], # pylint: disable=unused-argument
delta_token_ids: Sequence[int],
request: ChatCompletionRequest, # pylint: disable=unused-argument
) -> DeltaMessage | None:
start_in_text = self.tool_call_start_token in delta_text
start_in_ids = self.tool_call_start_token_id in delta_token_ids
tool_call_starting = start_in_text or start_in_ids
if tool_call_starting:
self._reset_streaming_state(tool_call_started=tool_call_starting)
self._tool_call_started_from_token_id = start_in_ids and not start_in_text
elif not previous_text:
if self._tool_call_started_from_token_id:
if current_text:
self._tool_call_started_from_token_id = False
else:
self._reset_streaming_state(tool_call_started=False)
if not self.is_tool_call_started:
return DeltaMessage(content=delta_text) if delta_text else None
content_before = None
if start_in_text:
before = delta_text[: delta_text.index(self.tool_call_start_token)]
content_before = before or None
delta_tool_call = self._extract_delta_tool_call(current_text)
if delta_tool_call:
return DeltaMessage(
content=content_before,
tool_calls=[delta_tool_call],
)
if content_before:
return DeltaMessage(content=content_before)
if (
not delta_text
and delta_token_ids
and self.prev_tool_call_arr
and self.tool_call_end_token_id not in delta_token_ids
):
return DeltaMessage(content="")
return None
MinimaxM2ToolParser.__init__ = _patched_init
MinimaxM2ToolParser._parse_single_invoke = _patched_parse_single_invoke
MinimaxM2ToolParser._reset_streaming_state = _reset_streaming_state
MinimaxM2ToolParser._ensure_streaming_slots = _ensure_streaming_slots
MinimaxM2ToolParser._get_param_config = _get_param_config
MinimaxM2ToolParser._serialize_partial_param_value = _serialize_partial_param_value
MinimaxM2ToolParser._build_partial_arguments = _build_partial_arguments
MinimaxM2ToolParser._get_invoke_states = _get_invoke_states
MinimaxM2ToolParser._finalize_completed_tool_call = _finalize_completed_tool_call
MinimaxM2ToolParser._extract_delta_tool_call = _extract_delta_tool_call
MinimaxM2ToolParser.extract_tool_calls_streaming = _patched_extract_tool_calls_streaming

View File

@@ -0,0 +1,462 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# MiniMax-M2 usage accounting: backport reasoning-token usage details.
#
from __future__ import annotations
import json
from collections.abc import AsyncIterator, Sequence
from dataclasses import dataclass
from types import MethodType
from typing import Any
from vllm.entrypoints.openai.chat_completion import protocol as chat_protocol
from vllm.entrypoints.openai.chat_completion import serving as chat_serving
from vllm.entrypoints.openai.chat_completion.serving import OpenAIServingChat
from vllm.entrypoints.openai.engine import protocol as engine_protocol
from vllm.reasoning import minimax_m2_reasoning_parser as minimax_parser
_MINIMAX_REASONING_PARSER_TYPES = (
minimax_parser.MiniMaxM2ReasoningParser,
minimax_parser.MiniMaxM2AppendThinkReasoningParser,
)
class CompletionTokenUsageInfo(engine_protocol.OpenAIBaseModel):
reasoning_tokens: int | None = None
audio_tokens: int | None = None
accepted_prediction_tokens: int | None = None
rejected_prediction_tokens: int | None = None
class UsageInfo(engine_protocol.UsageInfo):
completion_tokens_details: CompletionTokenUsageInfo | None = None
CompletionTokenUsageInfo.__module__ = engine_protocol.__name__
UsageInfo.__module__ = engine_protocol.__name__
# The OpenAI usage schema is process-wide. Keep only this schema backfill
# global; the expensive token tracking below is bound to MiniMax instances.
engine_protocol.CompletionTokenUsageInfo = CompletionTokenUsageInfo
engine_protocol.UsageInfo = UsageInfo
chat_protocol.UsageInfo = UsageInfo
chat_serving.CompletionTokenUsageInfo = CompletionTokenUsageInfo
chat_serving.UsageInfo = UsageInfo
def _rebuild_model_field(model_cls, field_name: str, annotation) -> None:
model_cls.__annotations__[field_name] = annotation
model_cls.model_fields[field_name].annotation = annotation
model_cls.model_rebuild(force=True)
_rebuild_model_field(chat_protocol.ChatCompletionResponse, "usage", UsageInfo)
_rebuild_model_field(chat_protocol.ChatCompletionStreamResponse, "usage", UsageInfo | None)
_rebuild_model_field(engine_protocol.RequestResponseMetadata, "final_usage_info", UsageInfo | None)
def _count_minimax_reasoning_tokens(
token_ids: Sequence[int],
end_token_id: int | None,
) -> int:
if end_token_id is None:
return 0
for idx, token_id in enumerate(token_ids):
if token_id == end_token_id:
return idx
return len(token_ids)
def _patched_count_reasoning_tokens(self, token_ids: Sequence[int]) -> int:
return _count_minimax_reasoning_tokens(token_ids, self.end_token_id)
minimax_parser.MiniMaxM2ReasoningParser.count_reasoning_tokens = _patched_count_reasoning_tokens
minimax_parser.MiniMaxM2AppendThinkReasoningParser.count_reasoning_tokens = _patched_count_reasoning_tokens
def _count_minimax_reasoning_tokens_for_usage(
token_ids: Sequence[int],
reasoning_parser,
) -> int | None:
reasoning_parser = _resolve_reasoning_parser(reasoning_parser)
if reasoning_parser is None or not _is_minimax_reasoning_parser(reasoning_parser):
return None
count_reasoning_tokens = getattr(reasoning_parser, "count_reasoning_tokens", None)
if count_reasoning_tokens is None:
return None
return count_reasoning_tokens(token_ids)
def _resolve_reasoning_parser(reasoning_parser):
if reasoning_parser is None:
return None
return getattr(reasoning_parser, "reasoning_parser", reasoning_parser)
def _is_minimax_reasoning_parser(reasoning_parser) -> bool:
return isinstance(
_resolve_reasoning_parser(reasoning_parser),
_MINIMAX_REASONING_PARSER_TYPES,
)
def _clamp_reasoning_tokens(
reasoning_tokens: int | None,
completion_tokens: int,
) -> int | None:
if reasoning_tokens is None:
return None
return max(0, min(reasoning_tokens, completion_tokens))
def _make_usage_info(
self,
*,
prompt_tokens: int,
completion_tokens: int,
num_cached_tokens: int | None = None,
reasoning_tokens: int | None = None,
) -> UsageInfo:
usage = UsageInfo(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
)
reasoning_tokens = _clamp_reasoning_tokens(reasoning_tokens, completion_tokens)
if reasoning_tokens is not None:
usage.completion_tokens_details = CompletionTokenUsageInfo(reasoning_tokens=reasoning_tokens)
if self.enable_prompt_tokens_details and num_cached_tokens is not None:
usage.prompt_tokens_details = chat_serving.PromptTokenUsageInfo(cached_tokens=num_cached_tokens)
return usage
def _is_minimax_reasoning_parser_cls(reasoning_parser_cls) -> bool:
return isinstance(reasoning_parser_cls, type) and issubclass(
reasoning_parser_cls,
_MINIMAX_REASONING_PARSER_TYPES,
)
@dataclass
class _UsageTrackingState:
completion_tokens: list[int]
raw_output_token_ids: list[list[int]]
reasoning_parser: Any
enable_prompt_tokens_details: bool = False
num_prompt_tokens: int = 0
num_cached_tokens: int | None = None
final_res: Any = None
def _create_usage_tracking_state(
num_choices: int,
reasoning_parser,
enable_prompt_tokens_details: bool = False,
) -> _UsageTrackingState:
return _UsageTrackingState(
completion_tokens=[0] * num_choices,
raw_output_token_ids=[[] for _ in range(num_choices)],
reasoning_parser=reasoning_parser,
enable_prompt_tokens_details=enable_prompt_tokens_details,
)
def _update_usage_tracking_state(
state: _UsageTrackingState,
res,
) -> None:
if res.prompt_token_ids is not None:
num_prompt_tokens = len(res.prompt_token_ids)
if res.encoder_prompt_token_ids is not None:
num_prompt_tokens += len(res.encoder_prompt_token_ids)
state.num_prompt_tokens = num_prompt_tokens
if state.num_cached_tokens is None:
state.num_cached_tokens = res.num_cached_tokens
state.final_res = res
for output in res.outputs:
if 0 <= output.index < len(state.completion_tokens):
token_ids = chat_serving.as_list(output.token_ids)
state.completion_tokens[output.index] += len(token_ids)
state.raw_output_token_ids[output.index].extend(token_ids)
async def _tracked_result_generator(
result_generator: AsyncIterator,
state: _UsageTrackingState,
):
async for res in result_generator:
_update_usage_tracking_state(state, res)
yield res
def _sum_reasoning_tokens_for_usage(
raw_output_token_ids: list[list[int]],
reasoning_parser,
) -> int | None:
if reasoning_parser is None:
return None
reasoning_token_counts = [
_count_minimax_reasoning_tokens_for_usage(token_ids, reasoning_parser) for token_ids in raw_output_token_ids
]
if all(reasoning_tokens is None for reasoning_tokens in reasoning_token_counts):
return None
return sum(reasoning_tokens or 0 for reasoning_tokens in reasoning_token_counts)
def _reasoning_tokens_for_choice(
state: _UsageTrackingState,
choice_index: int,
) -> int | None:
if state.reasoning_parser is None:
return None
if not 0 <= choice_index < len(state.raw_output_token_ids):
return None
return _count_minimax_reasoning_tokens_for_usage(
state.raw_output_token_ids[choice_index],
state.reasoning_parser,
)
def _make_full_response_usage(
self,
state: _UsageTrackingState,
) -> UsageInfo | None:
if state.final_res is None:
return None
return self._make_usage_info(
prompt_tokens=state.num_prompt_tokens,
completion_tokens=sum(state.completion_tokens),
num_cached_tokens=state.num_cached_tokens,
reasoning_tokens=_sum_reasoning_tokens_for_usage(
state.raw_output_token_ids,
state.reasoning_parser,
),
)
def _usage_reasoning_tokens_for_stream_chunk(
state: _UsageTrackingState,
chunk: dict[str, Any],
completion_tokens: int,
) -> int | None:
if state.reasoning_parser is None:
return None
choices = chunk.get("choices") or []
if choices:
choice_index = choices[0].get("index", 0)
reasoning_tokens = _reasoning_tokens_for_choice(state, choice_index)
else:
reasoning_tokens = _sum_reasoning_tokens_for_usage(
state.raw_output_token_ids,
state.reasoning_parser,
)
return _clamp_reasoning_tokens(reasoning_tokens, completion_tokens)
def _inject_stream_usage_details(
data: str,
state: _UsageTrackingState,
) -> str:
prefix = "data: "
suffix = "\n\n"
if not data.startswith(prefix):
return data
payload = data[len(prefix) :]
if payload.endswith(suffix):
payload = payload[: -len(suffix)]
if payload == "[DONE]":
return data
try:
chunk = json.loads(payload)
except json.JSONDecodeError:
return data
usage = chunk.get("usage")
if not isinstance(usage, dict):
return data
updated_usage = False
if state.enable_prompt_tokens_details and state.num_cached_tokens is not None:
usage["prompt_tokens_details"] = {
"cached_tokens": state.num_cached_tokens,
}
updated_usage = True
completion_tokens = usage.get("completion_tokens") or 0
reasoning_tokens = _usage_reasoning_tokens_for_stream_chunk(
state,
chunk,
completion_tokens,
)
if reasoning_tokens is not None:
usage["completion_tokens_details"] = {
"reasoning_tokens": reasoning_tokens,
}
updated_usage = True
if not updated_usage:
return data
return f"{prefix}{json.dumps(chunk, ensure_ascii=False)}{suffix}"
async def _wrapped_chat_completion_stream_generator(
self,
request: chat_protocol.ChatCompletionRequest,
result_generator: AsyncIterator,
request_id: str,
model_name: str,
conversation,
tokenizer,
request_metadata: engine_protocol.RequestResponseMetadata,
reasoning_parser=None,
**extra_kwargs: Any,
):
original_stream_generator = self._ascend_original_chat_completion_stream_generator
num_choices = 1 if request.n is None else request.n
state = _create_usage_tracking_state(
num_choices,
reasoning_parser,
enable_prompt_tokens_details=self.enable_prompt_tokens_details,
)
async for data in original_stream_generator(
request,
_tracked_result_generator(result_generator, state),
request_id,
model_name,
conversation,
tokenizer,
request_metadata,
reasoning_parser,
**extra_kwargs,
):
yield _inject_stream_usage_details(data, state)
usage = _make_full_response_usage(self, state)
if usage is not None:
request_metadata.final_usage_info = usage
async def _wrapped_chat_completion_full_generator(
self,
request: chat_protocol.ChatCompletionRequest,
result_generator: AsyncIterator,
request_id: str,
model_name: str,
conversation,
tokenizer,
request_metadata: engine_protocol.RequestResponseMetadata,
reasoning_parser=None,
):
original_full_generator = self._ascend_original_chat_completion_full_generator
num_choices = 1 if request.n is None else request.n
state = _create_usage_tracking_state(
num_choices,
reasoning_parser,
enable_prompt_tokens_details=self.enable_prompt_tokens_details,
)
response = await original_full_generator(
request,
_tracked_result_generator(result_generator, state),
request_id,
model_name,
conversation,
tokenizer,
request_metadata,
reasoning_parser,
)
if not isinstance(response, chat_protocol.ChatCompletionResponse):
return response
usage = _make_full_response_usage(self, state)
if usage is None:
return response
response.usage = usage
request_metadata.final_usage_info = usage
return response
_wrapped_chat_completion_stream_generator.__module__ = OpenAIServingChat.__module__
_wrapped_chat_completion_stream_generator.__qualname__ = (
f"{OpenAIServingChat.__qualname__}.chat_completion_stream_generator"
)
_wrapped_chat_completion_full_generator.__module__ = OpenAIServingChat.__module__
_wrapped_chat_completion_full_generator.__qualname__ = (
f"{OpenAIServingChat.__qualname__}.chat_completion_full_generator"
)
def _should_patch_chat_usage_instance(self) -> bool:
return _is_minimax_reasoning_parser_cls(self.reasoning_parser_cls)
def _patch_chat_usage_instance(self) -> None:
if getattr(self, "_ascend_minimax_usage_patched", False):
return
self._make_usage_info = MethodType(_make_usage_info, self)
self._ascend_original_chat_completion_stream_generator = MethodType(
OpenAIServingChat.chat_completion_stream_generator,
self,
)
self._ascend_original_chat_completion_full_generator = MethodType(
OpenAIServingChat.chat_completion_full_generator,
self,
)
self.chat_completion_stream_generator = MethodType(
_wrapped_chat_completion_stream_generator,
self,
)
self.chat_completion_full_generator = MethodType(
_wrapped_chat_completion_full_generator,
self,
)
self._ascend_minimax_usage_patched = True
class _ReasoningParserClsDescriptor:
def __init__(self, default_value=None):
self.default_value = default_value
def __get__(self, instance, owner=None):
if instance is None:
return self.default_value
return instance.__dict__.get("_ascend_reasoning_parser_cls", self.default_value)
def __set__(self, instance, value) -> None:
instance.__dict__["_ascend_reasoning_parser_cls"] = value
if _is_minimax_reasoning_parser_cls(value):
_patch_chat_usage_instance(instance)
_current_reasoning_parser_cls = OpenAIServingChat.__dict__.get("reasoning_parser_cls")
if not isinstance(_current_reasoning_parser_cls, _ReasoningParserClsDescriptor):
OpenAIServingChat.reasoning_parser_cls = _ReasoningParserClsDescriptor(_current_reasoning_parser_cls)

View File

@@ -0,0 +1,48 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# PR vllm-project/vllm#32623 introduced a new MLAPrefillBackend abstraction.
# When MLAAttention.__init__ calls get_mla_prefill_backend(), the upstream
# selector sees that Ascend NPU returns None for get_device_capability() and
# falls back to FlashAttnPrefillBackend, which asserts flash_attn_varlen_func
# is available — crashing on Ascend.
#
# Ascend's AscendSFAImpl/AscendMLAImpl handles the full forward pass (including
# prefill) via impl.forward(), so prefill_backend.run_prefill_* is never called.
# We register a no-op AscendMLAPrefillBackend and patch get_mla_prefill_backend
# so that MLAAttention.__init__ completes without error.
import torch
import vllm.model_executor.layers.attention.mla_attention
from vllm.v1.attention.backends.mla.prefill.base import MLAPrefillBackend
class AscendMLAPrefillBackend(MLAPrefillBackend):
@staticmethod
def get_name() -> str:
return "ASCEND"
@classmethod
def is_available(cls) -> bool:
return True
def run_prefill_new_tokens(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
return_softmax_lse: bool,
) -> torch.Tensor:
raise NotImplementedError("Ascend MLA prefill is handled by AscendSFAImpl/AscendMLAImpl")
def run_prefill_context_chunk(
self,
chunk_idx: int,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
raise NotImplementedError("Ascend MLA prefill is handled by AscendSFAImpl/AscendMLAImpl")
vllm.model_executor.layers.attention.mla_attention.get_mla_prefill_backend = lambda vllm_config: AscendMLAPrefillBackend

View File

@@ -0,0 +1,211 @@
from __future__ import annotations
import weakref
from collections import deque
from collections.abc import Callable
from multiprocessing.synchronize import Lock as LockType
import vllm.v1.executor.multiproc_executor
from vllm import envs
from vllm.config import VllmConfig
from vllm.distributed.device_communicators.shm_broadcast import Handle, MessageQueue
from vllm.utils.network_utils import get_distributed_init_method, get_loopback_ip, get_open_port
from vllm.utils.system_utils import get_mp_context
from vllm.v1.executor.abstract import FailureCallback
from vllm.v1.executor.multiproc_executor import (
FutureWrapper,
MultiprocExecutor,
UnreadyWorkerProcHandle,
WorkerProc,
set_multiprocessing_worker_envs,
)
class AscendMultiprocExecutor(MultiprocExecutor):
def _init_executor(self) -> None:
# Call self.shutdown at exit to clean up
# and ensure workers will be terminated.
self._finalizer = weakref.finalize(self, self.shutdown)
self.is_failed = False
self.failure_callback: FailureCallback | None = None
tensor_parallel_size, pp_parallel_size, pcp_parallel_size = self._get_parallel_sizes()
assert self.world_size == tensor_parallel_size * pp_parallel_size * pcp_parallel_size, (
f"world_size ({self.world_size}) must be equal to the "
f"tensor_parallel_size ({tensor_parallel_size}) x pipeline"
f"_parallel_size ({pp_parallel_size}) x prefill_context"
f"_parallel_size ({pcp_parallel_size}). "
)
# Set multiprocessing envs
set_multiprocessing_worker_envs()
# Multiprocessing-based executor does not support multi-node setting.
# Since it only works for single node, we can use the loopback address
# get_loopback_ip() for communication.
distributed_init_method = get_distributed_init_method(get_loopback_ip(), get_open_port())
self.rpc_broadcast_mq: MessageQueue | None = None
scheduler_output_handle: Handle | None = None
# Initialize worker and set up message queues for SchedulerOutputs
# and ModelRunnerOutputs
if self.parallel_config.node_rank_within_dp == 0:
# For leader node within each dp rank,
# each dp will have its own leader multiproc executor.
max_chunk_bytes = envs.VLLM_MQ_MAX_CHUNK_BYTES_MB * 1024 * 1024
self.rpc_broadcast_mq = MessageQueue(
self.world_size,
self.local_world_size,
max_chunk_bytes=max_chunk_bytes,
connect_ip=self.parallel_config.master_addr,
)
scheduler_output_handle = self.rpc_broadcast_mq.export_handle()
# Create workers
context = get_mp_context()
shared_worker_lock = context.Lock()
unready_workers: list[UnreadyWorkerProcHandle] = []
success = False
try:
global_start_rank = self.local_world_size * self.parallel_config.node_rank_within_dp
# When using fork, keep track of socket file descriptors that are
# inherited by the worker, so that we can close them in subsequent
# workers
inherited_fds: list[int] | None = [] if context.get_start_method() == "fork" else None
for local_rank in range(self.local_world_size):
global_rank = global_start_rank + local_rank
is_driver_worker = self._is_driver_worker(global_rank)
unready_worker_handle = AscendWorkerProc.make_worker_process(
vllm_config=self.vllm_config,
local_rank=local_rank,
rank=global_rank,
distributed_init_method=distributed_init_method,
input_shm_handle=scheduler_output_handle,
shared_worker_lock=shared_worker_lock,
is_driver_worker=is_driver_worker,
inherited_fds=inherited_fds,
)
unready_workers.append(unready_worker_handle)
if inherited_fds is not None:
inherited_fds.append(unready_worker_handle.death_writer.fileno())
inherited_fds.append(unready_worker_handle.ready_pipe.fileno())
# Workers must be created before wait_for_ready to avoid
# deadlock, since worker.init_device() does a device sync.
# Wait for all local workers to be ready.
self.workers = AscendWorkerProc.wait_for_ready(unready_workers)
# Start background thread to monitor worker health if not in headless mode.
if self.monitor_workers:
self.start_worker_monitor()
self.response_mqs = []
# Only leader node have remote response mqs
if self.parallel_config.node_rank_within_dp == 0:
for rank in range(self.world_size):
if rank < self.local_world_size:
local_message_queue = self.workers[rank].worker_response_mq
assert local_message_queue is not None
self.response_mqs.append(local_message_queue)
else:
remote_message_queue = self.workers[0].peer_worker_response_mqs[rank]
assert remote_message_queue is not None
self.response_mqs.append(remote_message_queue)
# Ensure message queues are ready. Will deadlock if re-ordered
# Must be kept consistent with the WorkerProc.
# Wait for all input mqs to be ready.
if self.rpc_broadcast_mq is not None:
self.rpc_broadcast_mq.wait_until_ready()
# Wait for all remote response mqs to be ready.
for response_mq in self.response_mqs:
response_mq.wait_until_ready()
self.futures_queue = deque[tuple[FutureWrapper, Callable]]()
self._post_init_executor()
success = True
finally:
if not success:
# Clean up the worker procs if there was a failure.
# Close death_writers first to signal workers to exit
for uw in unready_workers:
if uw.death_writer is not None:
uw.death_writer.close()
uw.death_writer = None
self._ensure_worker_termination([uw.proc for uw in unready_workers])
self.output_rank = self._get_output_rank()
def _get_parallel_sizes(self) -> tuple[int, int, int]:
self.world_size = self.parallel_config.world_size
assert self.world_size % self.parallel_config.nnodes_within_dp == 0, (
f"global world_size ({self.parallel_config.world_size}) must be "
f"divisible by nnodes_within_dp "
f"({self.parallel_config.nnodes_within_dp}). "
)
self.local_world_size = self.parallel_config.local_world_size
tp_size = self.parallel_config.tensor_parallel_size
pp_size = self.parallel_config.pipeline_parallel_size
pcp_size = self.parallel_config.prefill_context_parallel_size
return tp_size, pp_size, pcp_size
def _post_init_executor(self) -> None:
pass
def _is_driver_worker(self, rank: int) -> bool:
return rank % self.parallel_config.tensor_parallel_size == 0
class AscendWorkerProc(WorkerProc):
@staticmethod
def make_worker_process(
vllm_config: VllmConfig,
local_rank: int,
rank: int,
distributed_init_method: str,
input_shm_handle, # Receive SchedulerOutput
shared_worker_lock: LockType,
is_driver_worker: bool = False,
inherited_fds: list[int] | None = None,
) -> UnreadyWorkerProcHandle:
context = get_mp_context()
# Ready pipe to communicate readiness from child to parent
ready_reader, ready_writer = context.Pipe(duplex=False)
# Death pipe to let child detect parent process exit
death_reader, death_writer = context.Pipe(duplex=False)
if inherited_fds is not None:
inherited_fds = inherited_fds.copy()
inherited_fds.extend((ready_reader.fileno(), death_writer.fileno()))
process_kwargs = {
"vllm_config": vllm_config,
"local_rank": local_rank,
"rank": rank,
"distributed_init_method": distributed_init_method,
"input_shm_handle": input_shm_handle,
"ready_pipe": ready_writer,
"death_pipe": death_reader,
"shared_worker_lock": shared_worker_lock,
"is_driver_worker": is_driver_worker,
# Have the worker close parent end of this worker's pipes too
"inherited_fds": inherited_fds if inherited_fds is not None else [],
}
# Run EngineCore busy loop in background process.
proc = context.Process(
target=WorkerProc.worker_main,
kwargs=process_kwargs,
name=f"VllmWorker-{rank}",
daemon=False,
)
proc.start()
# Close child ends of pipes here in the parent
ready_writer.close()
death_reader.close()
# Keep death_writer open in parent - when parent exits,
# death_reader in child will get EOFError
return UnreadyWorkerProcHandle(proc, rank, ready_reader, death_writer)
vllm.v1.executor.multiproc_executor.MultiprocExecutor = AscendMultiprocExecutor

View File

@@ -0,0 +1,84 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
"""Backport vLLM PP + MTP runtime support.
The local Eagle/MTP drafter returns the draft tokens that belong to the model
output being processed. With PP batch_queue, EngineCore schedules a newer batch
before consuming the older output, so updating ``request.spec_token_ids`` from
``post_step`` observes live Request state from the newer schedule step.
"""
from __future__ import annotations
import copy
from functools import wraps
from vllm.logger import logger
_PATCHED = False
def _patch_model_config_validation() -> None:
from typing import get_args
from vllm.config.model import ModelConfig
from vllm.config.speculative import MTPModelTypes
original_verify = ModelConfig.verify_with_parallel_config
if getattr(original_verify, "_vllm_ascend_pp_mtp_patched", False):
return
mtp_model_types = set(get_args(MTPModelTypes))
@wraps(original_verify)
def _patched_verify_with_parallel_config(self, parallel_config):
hf_config = getattr(self, "hf_config", None)
model_type = getattr(hf_config, "model_type", None)
is_eagle_drafter = (model_type == "eagle" or model_type == "speculators") and any(
arch.startswith("Eagle") or arch.endswith("Eagle3") for arch in getattr(self, "architectures", ())
)
is_mtp_drafter = model_type in mtp_model_types
if (
getattr(self, "runner", None) == "draft"
and (is_eagle_drafter or is_mtp_drafter)
and getattr(parallel_config, "pipeline_parallel_size", 1) > 1
):
# Local Eagle/MTP drafters are loaded on the last PP stage rather
# than partitioned across all PP stages. Keep normal target-model
# validation intact, but validate these draft models as PP=1.
logger.warning(
"Validating local Eagle/MTP drafter with pipeline_parallel_size=1 "
"because it is loaded locally on the last pipeline stage."
)
patched_config = copy.copy(parallel_config)
patched_config.pipeline_parallel_size = 1
return original_verify(self, patched_config)
return original_verify(self, parallel_config)
_patched_verify_with_parallel_config._vllm_ascend_pp_mtp_patched = True # type: ignore[attr-defined]
ModelConfig.verify_with_parallel_config = _patched_verify_with_parallel_config
def _apply_patch() -> None:
global _PATCHED
if _PATCHED:
return
_PATCHED = True
_patch_model_config_validation()
_apply_patch()

View File

@@ -0,0 +1,234 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
"""Patches for profiling-based dynamic chunk sizing.
This module patches ``EngineCore`` to:
1. Run profiling at startup (after model_executor is ready).
2. Record execution timing after each model step to refine the
history-aware chunk prediction model online.
In multiprocessing ``spawn`` mode the child process starts a fresh Python
interpreter, so class-level monkey-patches applied in the parent are lost.
To handle this we additionally wrap ``EngineCoreProc.run_engine_core``
(the subprocess entry-point): when pickle resolves the wrapper it triggers
an import of this module, which re-applies the ``EngineCore.__init__``
patches inside the child process before any ``EngineCore`` is instantiated.
"""
from vllm.logger import logger
from vllm.v1.engine.core import EngineCore, EngineCoreProc
from vllm_ascend.utils import vllm_version_is
_profiling_patches_applied = False
_original_update_from_output = None
_original_schedule = None
# ---------------------------------------------------------------------------
# Helper: record execution timing
# ---------------------------------------------------------------------------
def _record_execution_timing(scheduler, scheduler_output, model_output):
"""Record execution timing for online model refinement.
Extracts ``execution_time_ms`` (set dynamically by the NPU model runner)
from the model output and feeds it back to the
``ProfilingChunkManager`` for incremental fitting of the history-aware
latency model.
"""
profiling_mgr = getattr(scheduler, "profiling_chunk_manager", None)
SET_TIME_COUNT = 3
if profiling_mgr is None or not profiling_mgr.is_ready:
return
# Once both the target latency and history model are calibrated,
# stop collecting timing data and disable the synchronize-and-time
# calls in the model runner to avoid unnecessary pipeline stalls.
if profiling_mgr._set_time_done and profiling_mgr.predictor.history_fitted:
try:
from vllm_ascend.ascend_config import get_ascend_config
get_ascend_config().profiling_chunk_config.need_timing = False
except RuntimeError:
pass
# Mark the scheduler so that the next scheduler_output carries
# a ``disable_profiling_timing`` flag to the worker process,
# which will set its own process-local need_timing to False.
scheduler._profiling_timing_done = True
return
elapsed_time_ms = getattr(model_output, "execution_time_ms", 0.0)
if elapsed_time_ms <= 0:
return
elapsed_time = elapsed_time_ms / 1000.0
try:
total_tokens = getattr(scheduler_output, "total_num_scheduled_tokens", 0)
if total_tokens <= 0:
return
num_scheduled_tokens = getattr(scheduler_output, "num_scheduled_tokens", {})
request_chunks = []
total_hist_tokens = 0
new_reqs = getattr(scheduler_output, "scheduled_new_reqs", [])
for req in new_reqs:
req_id = getattr(req, "request_id", None) or getattr(req, "req_id", None)
if req_id and req_id in num_scheduled_tokens:
chunk_size = num_scheduled_tokens[req_id]
hist_seq_len = getattr(req, "num_computed_tokens", 0)
total_hist_tokens += hist_seq_len
if chunk_size > 0:
request_chunks.append((chunk_size, hist_seq_len))
cached_reqs = getattr(scheduler_output, "scheduled_cached_reqs", None)
if cached_reqs is not None:
req_ids = getattr(cached_reqs, "req_ids", [])
computed_tokens_list = getattr(cached_reqs, "num_computed_tokens", [])
for i, req_id in enumerate(req_ids):
if req_id in num_scheduled_tokens:
chunk_size = num_scheduled_tokens[req_id]
hist_seq_len = computed_tokens_list[i] if i < len(computed_tokens_list) else 0
total_hist_tokens += hist_seq_len
if chunk_size > 0:
request_chunks.append((chunk_size, hist_seq_len))
# is first chunk processing — collect 3 samples before marking done
if total_hist_tokens == 0 and not profiling_mgr._set_time_done:
profiling_mgr.predictor.set_target_latency(0, elapsed_time * 1000)
profiling_mgr._set_time_count += 1
if profiling_mgr._set_time_count >= SET_TIME_COUNT:
profiling_mgr._set_time_done = True
if not request_chunks:
# Cannot accurately attribute batch latency to individual
# requests — skip this sample to avoid polluting the model.
logger.debug("[ProfilingChunk] Skipping timing sample: unable to extract per-request chunk info")
return
if not profiling_mgr.predictor.history_fitted:
profiling_mgr.record_batch_execution_time(request_chunks, elapsed_time)
except (AttributeError, TypeError) as e:
logger.debug("Failed to record execution timing: %s", e)
# ---------------------------------------------------------------------------
# Helper: wrap scheduler.update_from_output for timing
# ---------------------------------------------------------------------------
def _ensure_update_from_output_wrapped(scheduler):
"""Wrap scheduler.update_from_output to record execution timing."""
global _original_update_from_output
if _original_update_from_output is not None:
return
if not hasattr(scheduler, "profiling_chunk_manager"):
return
cls = type(scheduler)
_original_update_from_output = cls.update_from_output
def _wrapped_update_from_output(self, scheduler_output, model_output):
_record_execution_timing(self, scheduler_output, model_output)
return _original_update_from_output(self, scheduler_output, model_output)
cls.update_from_output = _wrapped_update_from_output
def _ensure_schedule_wrapped(scheduler):
"""Wrap scheduler.schedule to propagate timing-done signal via scheduler_output.
When ``_record_execution_timing`` detects that calibration is complete, it
sets ``scheduler._profiling_timing_done = True``. This wrapper copies that
flag onto every subsequent ``SchedulerOutput`` so the worker process can
read it and disable its own process-local ``need_timing``.
"""
global _original_schedule
if _original_schedule is not None:
return
if not hasattr(scheduler, "profiling_chunk_manager"):
return
cls = type(scheduler)
_original_schedule = cls.schedule
def _wrapped_schedule(self, throttle_prefills: bool = False):
if vllm_version_is("0.23.0"):
output = _original_schedule(self)
else:
output = _original_schedule(self, throttle_prefills)
if getattr(self, "_profiling_timing_done", False) and output is not None:
output.disable_profiling_timing = True
return output
cls.schedule = _wrapped_schedule
# ---------------------------------------------------------------------------
# Core: apply EngineCore.__init__ patches (idempotent)
# ---------------------------------------------------------------------------
def _apply_profiling_patches():
"""Patch ``EngineCore.__init__`` to trigger profiling and timing hooks.
Safe to call multiple times; the guard ``_profiling_patches_applied``
ensures the patch is applied at most once per process.
"""
global _profiling_patches_applied
if _profiling_patches_applied:
return
_profiling_patches_applied = True
original_init = EngineCore.__init__
def _patched_engine_core_init(self, *args, **kwargs):
original_init(self, *args, **kwargs)
if hasattr(self.scheduler, "run_profiling_chunk_init"):
logger.info("[ProfilingChunk] Running profiling initialization...")
self.scheduler.run_profiling_chunk_init(self.model_executor)
_ensure_update_from_output_wrapped(self.scheduler)
_ensure_schedule_wrapped(self.scheduler)
EngineCore.__init__ = _patched_engine_core_init
# ---------------------------------------------------------------------------
# 1. Apply patches at module level for the InprocClient (in-process) path.
# ---------------------------------------------------------------------------
_apply_profiling_patches()
# ---------------------------------------------------------------------------
# 2. Wrap EngineCoreProc.run_engine_core so that spawned subprocesses
# re-apply the patches. When the child unpickles this wrapper it
# imports this module, which triggers _apply_profiling_patches() above,
# ensuring EngineCore.__init__ is patched before any instance is created.
# ---------------------------------------------------------------------------
_original_run_engine_core = EngineCoreProc.run_engine_core
def _patched_run_engine_core(*args, **kwargs):
_apply_profiling_patches()
return _original_run_engine_core(*args, **kwargs)
EngineCoreProc.run_engine_core = _patched_run_engine_core

View File

@@ -0,0 +1,100 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import time
from contextlib import contextmanager
from vllm.distributed.device_communicators import shm_broadcast
MessageQueue = shm_broadcast.MessageQueue
# Cap on how long an idle reader parks before re-reading the authoritative SHM
# written-flag. Bounds lost-notify recovery latency to ~5s while the periodic
# wakeup stays negligible (one flag check per reader every 5s).
SHM_READER_RECHECK_INTERVAL_MS = 5000
def timeout_ms(self) -> int:
"""Returns a timeout, capped at the recheck interval, that is:
- min(time to deadline, time to next warning) if we're logging warnings
- time to deadline, if we're not logging warnings
- recheck interval if the timeout is None and we're not logging warnings
- raise TimeoutError if we are past the deadline
"""
wait_ms = SHM_READER_RECHECK_INTERVAL_MS
if self.warning_wait_time_ms is not None:
wait_ms = min(wait_ms, self.warning_wait_time_ms)
if self.timeout is None:
return wait_ms
time_left_ms = int((self.deadline - time.monotonic()) * 1000)
if time_left_ms <= 0:
raise TimeoutError
return min(wait_ms, time_left_ms)
@contextmanager
def acquire_read(
self,
timeout: float | None = None,
indefinite: bool = False,
):
assert self._is_local_reader, "Only readers can acquire read"
read_timeout = self.ReadTimeoutWithWarnings(timeout=timeout, should_warn=not indefinite)
with self.buffer.get_metadata(self.current_idx) as metadata_buffer:
while True:
def check():
shm_broadcast.memory_fence()
read_flag = metadata_buffer[self.local_reader_rank + 1]
written_flag = metadata_buffer[0]
return not (not written_flag or read_flag)
if shm_broadcast.SPINLOOP_EXT_ENABLED and not check():
shm_broadcast.spinloop(
metadata_buffer[0 : self.local_reader_rank + 1],
check,
timeout=shm_broadcast.SPINLOOP_TIMEOUT_SECONDS,
)
if not check():
# this block is either
# (1) not written
# (2) already read by this reader
# for readers, `self.current_idx` is the next block to read
# if this block is not ready,
# we need to wait until it is written
self._spin_condition.wait(timeout_ms=read_timeout.timeout_ms())
if self.shutting_down:
raise RuntimeError("cancelled")
# if we wait for a long time, log a message
if read_timeout.should_warn():
shm_broadcast.logger.info(
shm_broadcast.LONG_WAIT_TIME_LOG_MSG,
shm_broadcast.VLLM_RINGBUFFER_WARNING_INTERVAL,
)
continue
# found a block that is not read by this reader
# let caller read from the buffer
with self.buffer.get_data(self.current_idx) as buf:
try:
yield buf
finally:
# caller has read from the buffer; set the read flag.
metadata_buffer[self.local_reader_rank + 1] = 1
# Memory fence ensures the read flag is visible to the writer.
# Without this, writer may not see our read completion and
# could wait indefinitely for all readers to finish.
shm_broadcast.memory_fence()
next_idx = self.current_idx + 1
self.current_idx = next_idx % self.buffer.max_chunks
self._spin_condition.record_read()
break
MessageQueue.ReadTimeoutWithWarnings.timeout_ms = timeout_ms
MessageQueue.acquire_read = acquire_read

View File

@@ -0,0 +1,137 @@
from typing import TYPE_CHECKING, Any
from vllm.config.speculative import SpeculativeConfig
from vllm.utils.import_utils import LazyLoader
if TYPE_CHECKING:
import vllm.model_executor.layers.quantization as me_quant
from transformers import PretrainedConfig
else:
PretrainedConfig = Any
me_quant = LazyLoader("model_executor", globals(), "vllm.model_executor.layers.quantization")
def hf_config_override(hf_config: PretrainedConfig) -> PretrainedConfig:
initial_architecture = hf_config.architectures[0]
if hf_config.model_type in ("deepseek_v3", "deepseek_v32", "deepseek_v4", "glm_moe_dsa"):
target_model_type = hf_config.model_type
hf_config.model_type = "deepseek_mtp"
if hf_config.model_type == "deepseek_mtp":
if target_model_type == "deepseek_v4":
hf_config.update({"architectures": ["DeepSeekV4MTPModel"]})
else:
n_predict = getattr(hf_config, "num_nextn_predict_layers", None)
hf_config.update({"n_predict": n_predict, "architectures": ["DeepSeekMTPModel"]})
if hf_config.model_type in ("pangu_ultra_moe"):
hf_config.model_type = "pangu_ultra_moe_mtp"
if hf_config.model_type == "pangu_ultra_moe_mtp":
n_predict = getattr(hf_config, "num_nextn_predict_layers", None)
hf_config.update({"n_predict": n_predict, "architectures": ["OpenPanguMTPModel"]})
if hf_config.architectures[0] == "MiMoForCausalLM":
hf_config.model_type = "mimo_mtp"
n_predict = getattr(hf_config, "num_nextn_predict_layers", None)
hf_config.update(
{
"num_hidden_layers": 0,
"n_predict": n_predict,
"architectures": ["MiMoMTPModel"],
}
)
if hf_config.architectures[0] == "Glm4MoeForCausalLM":
hf_config.model_type = "glm4_moe_mtp"
n_predict = getattr(hf_config, "num_nextn_predict_layers", None)
hf_config.update(
{
"n_predict": n_predict,
"architectures": ["Glm4MoeMTPModel"],
}
)
if hf_config.architectures[0] == "Glm4MoeLiteForCausalLM":
hf_config.model_type = "glm4_moe_lite_mtp"
n_predict = getattr(hf_config, "num_nextn_predict_layers", None)
hf_config.update(
{
"num_hidden_layers": 0,
"n_predict": n_predict,
"architectures": ["Glm4MoeLiteMTPModel"],
}
)
if hf_config.architectures[0] == "GlmOcrForConditionalGeneration":
hf_config.model_type = "glm_ocr_mtp"
n_predict = getattr(hf_config, "num_nextn_predict_layers", None)
hf_config.update(
{
"num_hidden_layers": 0,
"n_predict": n_predict,
"architectures": ["GlmOcrMTPModel"],
}
)
if hf_config.model_type == "ernie4_5_moe":
hf_config.model_type = "ernie_mtp"
if hf_config.model_type == "ernie_mtp":
n_predict = getattr(hf_config, "num_nextn_predict_layers", None)
hf_config.update({"n_predict": n_predict, "architectures": ["ErnieMTPModel"]})
if (
hf_config.model_type == "nemotron_h"
and hasattr(hf_config, "num_nextn_predict_layers")
and hf_config.num_nextn_predict_layers > 0
):
# Check if this is an MTP variant
hf_config.model_type = "nemotron_h_mtp"
if hf_config.model_type == "nemotron_h_mtp":
n_predict = getattr(hf_config, "num_nextn_predict_layers", 1)
hf_config.update({"n_predict": n_predict, "architectures": ["NemotronHMTPModel"]})
if hf_config.model_type == "qwen3_next":
hf_config.model_type = "qwen3_next_mtp"
if hf_config.model_type == "qwen3_next_mtp":
n_predict = getattr(hf_config, "num_nextn_predict_layers", None)
hf_config.update({"n_predict": n_predict, "architectures": ["Qwen3NextMTP"]})
if hf_config.model_type == "exaone_moe":
hf_config.model_type = "exaone_moe_mtp"
if hf_config.model_type == "exaone_moe_mtp":
n_predict = getattr(hf_config, "num_nextn_predict_layers", None)
hf_config.update({"n_predict": n_predict, "architectures": ["ExaoneMoeMTP"]})
if hf_config.model_type in ("qwen3_5", "qwen3_5_moe"):
is_moe = hf_config.model_type == "qwen3_5_moe"
hf_config.model_type = "qwen3_5_mtp"
n_predict = getattr(hf_config, "mtp_num_hidden_layers", None)
hf_config.update(
{
"n_predict": n_predict,
"architectures": ["Qwen3_5MoeMTP" if is_moe else "Qwen3_5MTP"],
}
)
if hf_config.model_type == "longcat_flash":
hf_config.model_type = "longcat_flash_mtp"
n_predict = getattr(hf_config, "num_nextn_predict_layers", 1)
hf_config.update({"n_predict": n_predict, "architectures": ["LongCatFlashMTPModel"]})
if hf_config.model_type in ("step3p5", "step3p7") or hf_config.architectures[0] in (
"Step3p5ForCausalLM",
"Step3p7ForConditionalGeneration",
):
quantization_config = getattr(hf_config, "quantization_config", None)
hf_config = getattr(hf_config, "text_config", hf_config)
if quantization_config is not None and getattr(hf_config, "quantization_config", None) is None:
hf_config.update({"quantization_config": quantization_config})
hf_config.model_type = "step3p5_mtp"
n_predict = getattr(hf_config, "num_nextn_predict_layers", 1)
hf_config.update({"n_predict": n_predict, "architectures": ["Step3p5MTP"]})
if initial_architecture == "MistralLarge3ForCausalLM":
hf_config.update({"architectures": ["EagleMistralLarge3ForCausalLM"]})
return hf_config
SpeculativeConfig.hf_config_override = hf_config_override

View File

@@ -0,0 +1,134 @@
# SPDX-License-Identifier: Apache-2.0
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
from __future__ import annotations
from inspect import Signature, signature
from typing import Any
from vllm.exceptions import VLLMValidationError
from vllm.sampling_params import SamplingParams
from vllm.v1.structured_output import StructuredOutputManager
_BACKEND_ATTR = "_vllm_ascend_structured_output_backend"
_ORIGINAL_GRAMMAR_INIT_ATTR = "_vllm_ascend_original_grammar_init"
_ORIGINAL_VALIDATE_ATTR = "_vllm_ascend_original_validate_structured_outputs"
def _request_backend(request: Any) -> str | None:
if getattr(request, "structured_output_request", None) is None:
return None
sampling_params = getattr(request, "sampling_params", None)
structured_outputs = getattr(sampling_params, "structured_outputs", None)
backend = getattr(structured_outputs, "_backend", None)
return backend if isinstance(backend, str) else None
def _backend_name_from_instance(backend: Any) -> str | None:
if backend is None:
return None
backend_names = {
"XgrammarBackend": "xgrammar",
"GuidanceBackend": "guidance",
"OutlinesBackend": "outlines",
"LMFormatEnforcerBackend": "lm-format-enforcer",
}
for backend_cls in type(backend).__mro__:
for class_name, backend_name in backend_names.items():
if class_name in backend_cls.__name__:
return backend_name
return None
def _raise_mixed_backend(initialized_backend: str, request_backend: str) -> None:
raise VLLMValidationError(
"V1 structured outputs only supports one backend per engine. "
f"The engine is already using '{initialized_backend}', but "
f"this request resolved to '{request_backend}'. Configure "
"`structured_outputs_config.backend` explicitly or use schemas "
"supported by the initialized backend."
)
def _sampling_params_backend(sampling_params: SamplingParams) -> str | None:
structured_outputs = getattr(sampling_params, "structured_outputs", None)
backend = getattr(structured_outputs, "_backend", None)
return backend if isinstance(backend, str) else None
def _structured_outputs_config_from_call(
validate_signature: Signature,
sampling_params: SamplingParams,
args: tuple[Any, ...],
kwargs: dict[str, Any],
) -> Any:
bound_arguments = validate_signature.bind_partial(
sampling_params,
*args,
**kwargs,
)
return bound_arguments.arguments.get("structured_outputs_config")
def _patch_sampling_params_validation() -> None:
original_validate = SamplingParams._validate_structured_outputs
validate_signature = signature(original_validate)
setattr(SamplingParams, _ORIGINAL_VALIDATE_ATTR, original_validate)
def _validate_structured_outputs(
self: SamplingParams,
*args: Any,
**kwargs: Any,
) -> None:
result = original_validate(self, *args, **kwargs)
structured_outputs_config = _structured_outputs_config_from_call(
validate_signature,
self,
args,
kwargs,
)
request_backend = _sampling_params_backend(self)
if structured_outputs_config is None or request_backend is None:
return result
initialized_backend = getattr(structured_outputs_config, _BACKEND_ATTR, None)
if initialized_backend is not None and request_backend != initialized_backend:
_raise_mixed_backend(initialized_backend, request_backend)
setattr(structured_outputs_config, _BACKEND_ATTR, request_backend)
return result
SamplingParams._validate_structured_outputs = _validate_structured_outputs
def _patch_structured_output_manager() -> None:
original_grammar_init = StructuredOutputManager.grammar_init
setattr(StructuredOutputManager, _ORIGINAL_GRAMMAR_INIT_ATTR, original_grammar_init)
def grammar_init(self: StructuredOutputManager, request: Any) -> None:
request_backend = _request_backend(request)
if request_backend is None:
return original_grammar_init(self, request)
initialized_backend = getattr(self, _BACKEND_ATTR, None)
if initialized_backend is None:
initialized_backend = _backend_name_from_instance(getattr(self, "backend", None))
if initialized_backend is not None:
setattr(self, _BACKEND_ATTR, initialized_backend)
if initialized_backend is not None and request_backend != initialized_backend:
_raise_mixed_backend(initialized_backend, request_backend)
result = original_grammar_init(self, request)
if getattr(self, "backend", None) is not None:
setattr(self, _BACKEND_ATTR, request_backend)
return result
StructuredOutputManager.grammar_init = grammar_init
_patch_sampling_params_validation()
_patch_structured_output_manager()

View File

@@ -0,0 +1,87 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# OpenAI chat completions: omit empty tool_calls in serialized payloads.
#
from __future__ import annotations
import json
from typing import Any
from vllm.entrypoints.openai.chat_completion.protocol import (
ChatCompletionResponse,
ChatCompletionStreamResponse,
)
_original_chat_completion_response_model_dump = ChatCompletionResponse.model_dump
_original_chat_completion_stream_response_model_dump = ChatCompletionStreamResponse.model_dump
def _omit_empty_tool_calls(payload: Any) -> Any:
if not isinstance(payload, dict):
return payload
choices = payload.get("choices")
if not isinstance(choices, list):
return payload
for choice in choices:
if not isinstance(choice, dict):
continue
for field_name in ("message", "delta"):
message = choice.get(field_name)
if isinstance(message, dict) and message.get("tool_calls") == []:
message.pop("tool_calls")
return payload
def _patched_chat_completion_response_model_dump(self, *args, **kwargs):
return _omit_empty_tool_calls(_original_chat_completion_response_model_dump(self, *args, **kwargs))
def _dump_json(payload: Any, indent: int | None, ensure_ascii: bool) -> str:
separators = None if indent is not None else (",", ":")
return json.dumps(payload, ensure_ascii=ensure_ascii, indent=indent, separators=separators)
def _patched_chat_completion_response_model_dump_json(self, *args, **kwargs):
dump_kwargs = dict(kwargs)
indent = dump_kwargs.pop("indent", None)
ensure_ascii = dump_kwargs.pop("ensure_ascii", False)
dump_kwargs.setdefault("mode", "json")
payload = _patched_chat_completion_response_model_dump(self, *args, **dump_kwargs)
return _dump_json(payload, indent, ensure_ascii)
def _patched_chat_completion_stream_response_model_dump(self, *args, **kwargs):
return _omit_empty_tool_calls(_original_chat_completion_stream_response_model_dump(self, *args, **kwargs))
def _patched_chat_completion_stream_response_model_dump_json(self, *args, **kwargs):
dump_kwargs = dict(kwargs)
indent = dump_kwargs.pop("indent", None)
ensure_ascii = dump_kwargs.pop("ensure_ascii", False)
dump_kwargs.setdefault("mode", "json")
payload = _patched_chat_completion_stream_response_model_dump(self, *args, **dump_kwargs)
return _dump_json(payload, indent, ensure_ascii)
ChatCompletionResponse.model_dump = _patched_chat_completion_response_model_dump
ChatCompletionResponse.model_dump_json = _patched_chat_completion_response_model_dump_json
ChatCompletionStreamResponse.model_dump = _patched_chat_completion_stream_response_model_dump
ChatCompletionStreamResponse.model_dump_json = _patched_chat_completion_stream_response_model_dump_json

View File

@@ -0,0 +1,16 @@
import torch
def patch_empty_cache() -> None:
torch.npu.empty_cache()
torch.accelerator.empty_cache = patch_empty_cache
# Monkey-patch torch.accelerator memory APIs for NPU compatibility.
# Upstream vLLM (commit 747b068) replaced current_platform.memory_stats()
# with torch.accelerator.memory_stats(), but torch.accelerator does not
# properly delegate to NPU. We redirect to torch.npu.* equivalents.
torch.accelerator.memory_stats = torch.npu.memory_stats # type: ignore[attr-defined]
torch.accelerator.memory_reserved = torch.npu.memory_reserved # type: ignore[attr-defined]
torch.accelerator.reset_peak_memory_stats = torch.npu.reset_peak_memory_stats # type: ignore[attr-defined]

View File

@@ -0,0 +1,20 @@
import vllm.envs as envs
from vllm.config.vllm import VllmConfig
def _patched_use_v2_model_runner(self) -> bool:
"""Return VLLM_USE_V2_MODEL_RUNNER env directly.
The upstream use_v2_model_runner gate-keeps the v2 runner with
per-model architecture whitelists, Triton availability checks, and
feature-support inspections. On Ascend the v2 runner is controlled
purely by the VLLM_USE_V2_MODEL_RUNNER environment variable;
model-compatibility decisions are deferred to the NPU runner itself.
"""
use_v2 = envs.VLLM_USE_V2_MODEL_RUNNER
if use_v2 is not None:
return use_v2
return False
VllmConfig.use_v2_model_runner = property(_patched_use_v2_model_runner)

View File

@@ -0,0 +1,73 @@
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# 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.
#
# Patch target: vllm.distributed.weight_transfer.factory.WeightTransferEngineFactory
#
# Replace the "nccl" and "ipc" factory entries with Ascend equivalents so that
# --weight-transfer-config '{"backend": "nccl"}' loads HCCLWeightTransferEngine
# and '{"backend": "ipc"}' loads NPUIPCWeightTransferEngine instead of the
# (unavailable) NCCL / CUDA IPC engines on Ascend NPU.
#
# Why this approach (factory swap) instead of patching Literal["nccl", "ipc"]:
# WeightTransferConfig.backend is a pydantic Literal["nccl", "ipc"].
# Adding "hccl" / "npu_ipc" would require modifying pydantic core schemas —
# fragile across pydantic versions. Swapping the factory entries means users
# pass the already-accepted "nccl" / "ipc" strings, but the factory resolves
# them to HCCL / NPU IPC.
#
# Timing — guaranteed to run before first factory usage:
#
# vllm serve main()
# line 24: from vllm.entrypoints.utils import ...
# → vllm.platforms.__getattr__("current_platform")
# → resolve_current_platform_cls_qualname()
# → vllm_ascend:register() → NPUPlatform()
# → NPUPlatform.pre_register_and_update()
# → adapt_patch(is_global_patch=True)
# → imports vllm_ascend.patch.platform
# → THIS PATCH RUNS ← "nccl" now points to HCCLWeightTransferEngine
# ...
# lines 82-86: subparser_init() → make_arg_parser()
# line 87: parse_args() → validates backend="nccl" via Literal (passes)
# ...
# later: worker init → WeightTransferEngineFactory.create_engine(config)
# → config.backend == "nccl" → factory loads HCCLWeightTransferEngine
#
# Future Plan:
# Remove this patch when upstream vllm relaxes the Literal type to str
# or provides an extension point for out-of-tree backends.
from typing import TYPE_CHECKING
from vllm.distributed.weight_transfer.factory import WeightTransferEngineFactory
from vllm_ascend.distributed.weight_transfer.hccl_engine import (
HCCLWeightTransferEngine,
)
if TYPE_CHECKING:
from vllm.distributed.weight_transfer.base import WeightTransferEngine
def _load_npu_ipc_engine() -> "type[WeightTransferEngine]":
from vllm_ascend.distributed.weight_transfer.npu_ipc_engine import (
NPUIPCWeightTransferEngine,
)
return NPUIPCWeightTransferEngine
WeightTransferEngineFactory._registry["nccl"] = lambda: HCCLWeightTransferEngine
WeightTransferEngineFactory._registry["ipc"] = _load_npu_ipc_engine