@@ -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
|
||||
|
||||
190
vllm_ascend/patch/platform/patch_async_swa_kv_lifetime.py
Normal file
190
vllm_ascend/patch/platform/patch_async_swa_kv_lifetime.py
Normal 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
|
||||
745
vllm_ascend/patch/platform/patch_balance_schedule.py
Normal file
745
vllm_ascend/patch/platform/patch_balance_schedule.py
Normal 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
|
||||
28
vllm_ascend/patch/platform/patch_camem_allocator.py
Normal file
28
vllm_ascend/patch/platform/patch_camem_allocator.py
Normal 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
|
||||
802
vllm_ascend/patch/platform/patch_deepseek_v4_tool_call_parser.py
Normal file
802
vllm_ascend/patch/platform/patch_deepseek_v4_tool_call_parser.py
Normal 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
|
||||
89
vllm_ascend/patch/platform/patch_distributed.py
Normal file
89
vllm_ascend/patch/platform/patch_distributed.py
Normal 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()
|
||||
72
vllm_ascend/patch/platform/patch_dp_device_ids.py
Normal file
72
vllm_ascend/patch/platform/patch_dp_device_ids.py
Normal 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
|
||||
57
vllm_ascend/patch/platform/patch_fused_moe.py
Normal file
57
vllm_ascend/patch/platform/patch_fused_moe.py
Normal 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
|
||||
47
vllm_ascend/patch/platform/patch_glm47_tool_call_parser.py
Normal file
47
vllm_ascend/patch/platform/patch_glm47_tool_call_parser.py
Normal 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
|
||||
145
vllm_ascend/patch/platform/patch_glm_tool_call_streaming.py
Normal file
145
vllm_ascend/patch/platform/patch_glm_tool_call_streaming.py
Normal 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
|
||||
526
vllm_ascend/patch/platform/patch_kv_cache_coordinator.py
Normal file
526
vllm_ascend/patch/platform/patch_kv_cache_coordinator.py
Normal 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]
|
||||
359
vllm_ascend/patch/platform/patch_kv_cache_utils.py
Normal file
359
vllm_ascend/patch/platform/patch_kv_cache_utils.py
Normal 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
|
||||
149
vllm_ascend/patch/platform/patch_mamba_config.py
Normal file
149
vllm_ascend/patch/platform/patch_mamba_config.py
Normal 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
|
||||
103
vllm_ascend/patch/platform/patch_mamba_config_310.py
Normal file
103
vllm_ascend/patch/platform/patch_mamba_config_310.py
Normal 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
|
||||
84
vllm_ascend/patch/platform/patch_mamba_manager.py
Normal file
84
vllm_ascend/patch/platform/patch_mamba_manager.py
Normal 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
|
||||
136
vllm_ascend/patch/platform/patch_minimax_m2_config.py
Normal file
136
vllm_ascend/patch/platform/patch_minimax_m2_config.py
Normal 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
|
||||
490
vllm_ascend/patch/platform/patch_minimax_m2_tool_call_parser.py
Normal file
490
vllm_ascend/patch/platform/patch_minimax_m2_tool_call_parser.py
Normal 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
|
||||
462
vllm_ascend/patch/platform/patch_minimax_usage_accounting.py
Normal file
462
vllm_ascend/patch/platform/patch_minimax_usage_accounting.py
Normal 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)
|
||||
48
vllm_ascend/patch/platform/patch_mla_prefill_backend.py
Normal file
48
vllm_ascend/patch/platform/patch_mla_prefill_backend.py
Normal 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
|
||||
211
vllm_ascend/patch/platform/patch_multiproc_executor.py
Normal file
211
vllm_ascend/patch/platform/patch_multiproc_executor.py
Normal 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
|
||||
84
vllm_ascend/patch/platform/patch_pp_mtp.py
Normal file
84
vllm_ascend/patch/platform/patch_pp_mtp.py
Normal 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()
|
||||
234
vllm_ascend/patch/platform/patch_profiling_chunk.py
Normal file
234
vllm_ascend/patch/platform/patch_profiling_chunk.py
Normal 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
|
||||
100
vllm_ascend/patch/platform/patch_shm_broadcast.py
Normal file
100
vllm_ascend/patch/platform/patch_shm_broadcast.py
Normal 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
|
||||
137
vllm_ascend/patch/platform/patch_speculative_config.py
Normal file
137
vllm_ascend/patch/platform/patch_speculative_config.py
Normal 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
|
||||
134
vllm_ascend/patch/platform/patch_structured_output.py
Normal file
134
vllm_ascend/patch/platform/patch_structured_output.py
Normal 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()
|
||||
87
vllm_ascend/patch/platform/patch_tool_choice_none_content.py
Normal file
87
vllm_ascend/patch/platform/patch_tool_choice_none_content.py
Normal 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
|
||||
16
vllm_ascend/patch/platform/patch_torch_accelerator.py
Normal file
16
vllm_ascend/patch/platform/patch_torch_accelerator.py
Normal 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]
|
||||
20
vllm_ascend/patch/platform/patch_use_v2_model_runner.py
Normal file
20
vllm_ascend/patch/platform/patch_use_v2_model_runner.py
Normal 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)
|
||||
73
vllm_ascend/patch/platform/patch_weight_transfer_engine.py
Normal file
73
vllm_ascend/patch/platform/patch_weight_transfer_engine.py
Normal 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
|
||||
Reference in New Issue
Block a user