fix(build): strip \r\n from all .py files — CRLF breaks patch_ops.sh text matching

31 files had Windows line endings (\r\n) from merge commit. This causes
patch_ops.sh replace_once() to fail: anchor strings use \n but file
content has \r\n, so no match → patch fails → docker build fails.

Also added .gitattributes to force LF for all text files going forward.
This commit is contained in:
Claude
2026-08-14 01:06:49 +00:00
parent aa4b4992d1
commit 872be0effa
32 changed files with 13814 additions and 13804 deletions

10
.gitattributes vendored Normal file
View File

@@ -0,0 +1,10 @@
# Force LF line endings for all text files
* text=auto eol=lf
*.py text eol=lf
*.sh text eol=lf
*.cu text eol=lf
*.cuh text eol=lf
*.yaml text eol=lf
*.yml text eol=lf
*.md text eol=lf
Dockerfile text eol=lf

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -1,261 +1,261 @@
""" """
This file contains the command line arguments for the vLLM's This file contains the command line arguments for the vLLM's
OpenAI-compatible server. It is kept in a separate file for documentation OpenAI-compatible server. It is kept in a separate file for documentation
purposes. purposes.
""" """
import argparse import argparse
import json import json
import ssl import ssl
from typing import List, Optional, Sequence, Union from typing import List, Optional, Sequence, Union
from vllm.engine.arg_utils import AsyncEngineArgs, nullable_str from vllm.engine.arg_utils import AsyncEngineArgs, nullable_str
from vllm.entrypoints.chat_utils import validate_chat_template from vllm.entrypoints.chat_utils import validate_chat_template
from vllm.entrypoints.openai.serving_engine import (LoRAModulePath, from vllm.entrypoints.openai.serving_engine import (LoRAModulePath,
PromptAdapterPath) PromptAdapterPath)
from vllm.entrypoints.openai.tool_parsers import ToolParserManager from vllm.entrypoints.openai.tool_parsers import ToolParserManager
from vllm.utils import FlexibleArgumentParser from vllm.utils import FlexibleArgumentParser
class LoRAParserAction(argparse.Action): class LoRAParserAction(argparse.Action):
def __call__( def __call__(
self, self,
parser: argparse.ArgumentParser, parser: argparse.ArgumentParser,
namespace: argparse.Namespace, namespace: argparse.Namespace,
values: Optional[Union[str, Sequence[str]]], values: Optional[Union[str, Sequence[str]]],
option_string: Optional[str] = None, option_string: Optional[str] = None,
): ):
if values is None: if values is None:
values = [] values = []
if isinstance(values, str): if isinstance(values, str):
raise TypeError("Expected values to be a list") raise TypeError("Expected values to be a list")
lora_list: List[LoRAModulePath] = [] lora_list: List[LoRAModulePath] = []
for item in values: for item in values:
if item in [None, '']: # Skip if item is None or empty string if item in [None, '']: # Skip if item is None or empty string
continue continue
if '=' in item and ',' not in item: # Old format: name=path if '=' in item and ',' not in item: # Old format: name=path
name, path = item.split('=') name, path = item.split('=')
lora_list.append(LoRAModulePath(name, path)) lora_list.append(LoRAModulePath(name, path))
else: # Assume JSON format else: # Assume JSON format
try: try:
lora_dict = json.loads(item) lora_dict = json.loads(item)
lora = LoRAModulePath(**lora_dict) lora = LoRAModulePath(**lora_dict)
lora_list.append(lora) lora_list.append(lora)
except json.JSONDecodeError: except json.JSONDecodeError:
parser.error( parser.error(
f"Invalid JSON format for --lora-modules: {item}") f"Invalid JSON format for --lora-modules: {item}")
except TypeError as e: except TypeError as e:
parser.error( parser.error(
f"Invalid fields for --lora-modules: {item} - {str(e)}" f"Invalid fields for --lora-modules: {item} - {str(e)}"
) )
setattr(namespace, self.dest, lora_list) setattr(namespace, self.dest, lora_list)
class PromptAdapterParserAction(argparse.Action): class PromptAdapterParserAction(argparse.Action):
def __call__( def __call__(
self, self,
parser: argparse.ArgumentParser, parser: argparse.ArgumentParser,
namespace: argparse.Namespace, namespace: argparse.Namespace,
values: Optional[Union[str, Sequence[str]]], values: Optional[Union[str, Sequence[str]]],
option_string: Optional[str] = None, option_string: Optional[str] = None,
): ):
if values is None: if values is None:
values = [] values = []
if isinstance(values, str): if isinstance(values, str):
raise TypeError("Expected values to be a list") raise TypeError("Expected values to be a list")
adapter_list: List[PromptAdapterPath] = [] adapter_list: List[PromptAdapterPath] = []
for item in values: for item in values:
name, path = item.split('=') name, path = item.split('=')
adapter_list.append(PromptAdapterPath(name, path)) adapter_list.append(PromptAdapterPath(name, path))
setattr(namespace, self.dest, adapter_list) setattr(namespace, self.dest, adapter_list)
def make_arg_parser(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: def make_arg_parser(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
parser.add_argument("--host", parser.add_argument("--host",
type=nullable_str, type=nullable_str,
default=None, default=None,
help="host name") help="host name")
parser.add_argument("--port", type=int, default=8000, help="port number") parser.add_argument("--port", type=int, default=8000, help="port number")
parser.add_argument( parser.add_argument(
"--uvicorn-log-level", "--uvicorn-log-level",
type=str, type=str,
default="info", default="info",
choices=['debug', 'info', 'warning', 'error', 'critical', 'trace'], choices=['debug', 'info', 'warning', 'error', 'critical', 'trace'],
help="log level for uvicorn") help="log level for uvicorn")
parser.add_argument("--allow-credentials", parser.add_argument("--allow-credentials",
action="store_true", action="store_true",
help="allow credentials") help="allow credentials")
parser.add_argument("--allowed-origins", parser.add_argument("--allowed-origins",
type=json.loads, type=json.loads,
default=["*"], default=["*"],
help="allowed origins") help="allowed origins")
parser.add_argument("--allowed-methods", parser.add_argument("--allowed-methods",
type=json.loads, type=json.loads,
default=["*"], default=["*"],
help="allowed methods") help="allowed methods")
parser.add_argument("--allowed-headers", parser.add_argument("--allowed-headers",
type=json.loads, type=json.loads,
default=["*"], default=["*"],
help="allowed headers") help="allowed headers")
parser.add_argument("--api-key", parser.add_argument("--api-key",
type=nullable_str, type=nullable_str,
default=None, default=None,
help="If provided, the server will require this key " help="If provided, the server will require this key "
"to be presented in the header.") "to be presented in the header.")
parser.add_argument( parser.add_argument(
"--lora-modules", "--lora-modules",
type=nullable_str, type=nullable_str,
default=None, default=None,
nargs='+', nargs='+',
action=LoRAParserAction, action=LoRAParserAction,
help="LoRA module configurations in either 'name=path' format" help="LoRA module configurations in either 'name=path' format"
"or JSON format. " "or JSON format. "
"Example (old format): 'name=path' " "Example (old format): 'name=path' "
"Example (new format): " "Example (new format): "
"'{\"name\": \"name\", \"local_path\": \"path\", " "'{\"name\": \"name\", \"local_path\": \"path\", "
"\"base_model_name\": \"id\"}'") "\"base_model_name\": \"id\"}'")
parser.add_argument( parser.add_argument(
"--prompt-adapters", "--prompt-adapters",
type=nullable_str, type=nullable_str,
default=None, default=None,
nargs='+', nargs='+',
action=PromptAdapterParserAction, action=PromptAdapterParserAction,
help="Prompt adapter configurations in the format name=path. " help="Prompt adapter configurations in the format name=path. "
"Multiple adapters can be specified.") "Multiple adapters can be specified.")
parser.add_argument("--chat-template", parser.add_argument("--chat-template",
type=nullable_str, type=nullable_str,
default=None, default=None,
help="The file path to the chat template, " help="The file path to the chat template, "
"or the template in single-line form " "or the template in single-line form "
"for the specified model") "for the specified model")
parser.add_argument("--response-role", parser.add_argument("--response-role",
type=nullable_str, type=nullable_str,
default="assistant", default="assistant",
help="The role name to return if " help="The role name to return if "
"`request.add_generation_prompt=true`.") "`request.add_generation_prompt=true`.")
parser.add_argument("--ssl-keyfile", parser.add_argument("--ssl-keyfile",
type=nullable_str, type=nullable_str,
default=None, default=None,
help="The file path to the SSL key file") help="The file path to the SSL key file")
parser.add_argument("--ssl-certfile", parser.add_argument("--ssl-certfile",
type=nullable_str, type=nullable_str,
default=None, default=None,
help="The file path to the SSL cert file") help="The file path to the SSL cert file")
parser.add_argument("--ssl-ca-certs", parser.add_argument("--ssl-ca-certs",
type=nullable_str, type=nullable_str,
default=None, default=None,
help="The CA certificates file") help="The CA certificates file")
parser.add_argument( parser.add_argument(
"--ssl-cert-reqs", "--ssl-cert-reqs",
type=int, type=int,
default=int(ssl.CERT_NONE), default=int(ssl.CERT_NONE),
help="Whether client certificate is required (see stdlib ssl module's)" help="Whether client certificate is required (see stdlib ssl module's)"
) )
parser.add_argument( parser.add_argument(
"--root-path", "--root-path",
type=nullable_str, type=nullable_str,
default=None, default=None,
help="FastAPI root_path when app is behind a path based routing proxy") help="FastAPI root_path when app is behind a path based routing proxy")
parser.add_argument( parser.add_argument(
"--middleware", "--middleware",
type=nullable_str, type=nullable_str,
action="append", action="append",
default=[], default=[],
help="Additional ASGI middleware to apply to the app. " help="Additional ASGI middleware to apply to the app. "
"We accept multiple --middleware arguments. " "We accept multiple --middleware arguments. "
"The value should be an import path. " "The value should be an import path. "
"If a function is provided, vLLM will add it to the server " "If a function is provided, vLLM will add it to the server "
"using @app.middleware('http'). " "using @app.middleware('http'). "
"If a class is provided, vLLM will add it to the server " "If a class is provided, vLLM will add it to the server "
"using app.add_middleware(). ") "using app.add_middleware(). ")
parser.add_argument( parser.add_argument(
"--return-tokens-as-token-ids", "--return-tokens-as-token-ids",
action="store_true", action="store_true",
help="When --max-logprobs is specified, represents single tokens as " help="When --max-logprobs is specified, represents single tokens as "
"strings of the form 'token_id:{token_id}' so that tokens that " "strings of the form 'token_id:{token_id}' so that tokens that "
"are not JSON-encodable can be identified.") "are not JSON-encodable can be identified.")
parser.add_argument( parser.add_argument(
"--disable-frontend-multiprocessing", "--disable-frontend-multiprocessing",
action="store_true", action="store_true",
help="If specified, will run the OpenAI frontend server in the same " help="If specified, will run the OpenAI frontend server in the same "
"process as the model serving engine.") "process as the model serving engine.")
parser.add_argument( parser.add_argument(
"--enable-auto-tool-choice", "--enable-auto-tool-choice",
action="store_true", action="store_true",
default=False, default=False,
help= help=
"Enable auto tool choice for supported models. Use --tool-call-parser" "Enable auto tool choice for supported models. Use --tool-call-parser"
"to specify which parser to use") "to specify which parser to use")
valid_tool_parsers = ToolParserManager.tool_parsers.keys() valid_tool_parsers = ToolParserManager.tool_parsers.keys()
parser.add_argument( parser.add_argument(
"--tool-call-parser", "--tool-call-parser",
type=str, type=str,
metavar="{" + ",".join(valid_tool_parsers) + "} or name registered in " metavar="{" + ",".join(valid_tool_parsers) + "} or name registered in "
"--tool-parser-plugin", "--tool-parser-plugin",
default=None, default=None,
help= help=
"Select the tool call parser depending on the model that you're using." "Select the tool call parser depending on the model that you're using."
" This is used to parse the model-generated tool call into OpenAI API " " This is used to parse the model-generated tool call into OpenAI API "
"format. Required for --enable-auto-tool-choice.") "format. Required for --enable-auto-tool-choice.")
parser.add_argument( parser.add_argument(
"--tool-parser-plugin", "--tool-parser-plugin",
type=str, type=str,
default="", default="",
help= help=
"Special the tool parser plugin write to parse the model-generated tool" "Special the tool parser plugin write to parse the model-generated tool"
" into OpenAI API format, the name register in this plugin can be used " " into OpenAI API format, the name register in this plugin can be used "
"in --tool-call-parser.") "in --tool-call-parser.")
parser.add_argument( parser.add_argument(
"--reasoning-parser", "--reasoning-parser",
type=str, type=str,
default=None, default=None,
help= help=
"Select the reasoning parser to split <think>...</think> content into " "Select the reasoning parser to split <think>...</think> content into "
"reasoning_content vs content in the response. " "reasoning_content vs content in the response. "
"Supported: qwen3") "Supported: qwen3")
parser = AsyncEngineArgs.add_cli_args(parser) parser = AsyncEngineArgs.add_cli_args(parser)
parser.add_argument('--max-log-len', parser.add_argument('--max-log-len',
type=int, type=int,
default=None, default=None,
help='Max number of prompt characters or prompt ' help='Max number of prompt characters or prompt '
'ID numbers being printed in log.' 'ID numbers being printed in log.'
'\n\nDefault: Unlimited') '\n\nDefault: Unlimited')
parser.add_argument( parser.add_argument(
"--disable-fastapi-docs", "--disable-fastapi-docs",
action='store_true', action='store_true',
default=False, default=False,
help="Disable FastAPI's OpenAPI schema, Swagger UI, and ReDoc endpoint" help="Disable FastAPI's OpenAPI schema, Swagger UI, and ReDoc endpoint"
) )
return parser return parser
def validate_parsed_serve_args(args: argparse.Namespace): def validate_parsed_serve_args(args: argparse.Namespace):
"""Quick checks for model serve args that raise prior to loading.""" """Quick checks for model serve args that raise prior to loading."""
if hasattr(args, "subparser") and args.subparser != "serve": if hasattr(args, "subparser") and args.subparser != "serve":
return return
# Ensure that the chat template is valid; raises if it likely isn't # Ensure that the chat template is valid; raises if it likely isn't
validate_chat_template(args.chat_template) validate_chat_template(args.chat_template)
# Enable auto tool needs a tool call parser to be valid # Enable auto tool needs a tool call parser to be valid
if args.enable_auto_tool_choice and not args.tool_call_parser: if args.enable_auto_tool_choice and not args.tool_call_parser:
raise TypeError("Error: --enable-auto-tool-choice requires " raise TypeError("Error: --enable-auto-tool-choice requires "
"--tool-call-parser") "--tool-call-parser")
def create_parser_for_docs() -> FlexibleArgumentParser: def create_parser_for_docs() -> FlexibleArgumentParser:
parser_for_docs = FlexibleArgumentParser( parser_for_docs = FlexibleArgumentParser(
prog="-m vllm.entrypoints.openai.api_server") prog="-m vllm.entrypoints.openai.api_server")
return make_arg_parser(parser_for_docs) return make_arg_parser(parser_for_docs)

View File

@@ -1,224 +1,224 @@
from typing import Dict, List, Optional from typing import Dict, List, Optional
import torch import torch
from vllm.attention.backends.abstract import AttentionMetadata from vllm.attention.backends.abstract import AttentionMetadata
class MambaCacheManager: class MambaCacheManager:
def __init__(self, dtype, num_mamba_layers, max_batch_size, def __init__(self, dtype, num_mamba_layers, max_batch_size,
conv_state_shape, temporal_state_shape): conv_state_shape, temporal_state_shape):
conv_state = torch.empty(size=(num_mamba_layers, max_batch_size) + conv_state = torch.empty(size=(num_mamba_layers, max_batch_size) +
conv_state_shape, conv_state_shape,
dtype=dtype, dtype=dtype,
device="cuda") device="cuda")
temporal_state = torch.zeros(size=(num_mamba_layers, max_batch_size) + temporal_state = torch.zeros(size=(num_mamba_layers, max_batch_size) +
temporal_state_shape, temporal_state_shape,
dtype=dtype, dtype=dtype,
device="cuda") device="cuda")
self.mamba_cache = (conv_state, temporal_state) self.mamba_cache = (conv_state, temporal_state)
# Maps between the request id and a dict that maps between the seq_id # Maps between the request id and a dict that maps between the seq_id
# and its index inside the self.mamba_cache # and its index inside the self.mamba_cache
self.mamba_cache_indices_mapping: Dict[str, Dict[int, int]] = {} self.mamba_cache_indices_mapping: Dict[str, Dict[int, int]] = {}
def current_run_tensors(self, input_ids: torch.Tensor, def current_run_tensors(self, input_ids: torch.Tensor,
attn_metadata: AttentionMetadata, **kwargs): attn_metadata: AttentionMetadata, **kwargs):
""" """
Return the tensors for the current run's conv and ssm state. Return the tensors for the current run's conv and ssm state.
""" """
if "seqlen_agnostic_capture_inputs" not in kwargs: if "seqlen_agnostic_capture_inputs" not in kwargs:
# We get here only on Prefill/Eager mode runs # We get here only on Prefill/Eager mode runs
request_ids_to_seq_ids = kwargs["request_ids_to_seq_ids"] request_ids_to_seq_ids = kwargs["request_ids_to_seq_ids"]
finished_requests_ids = kwargs["finished_requests_ids"] finished_requests_ids = kwargs["finished_requests_ids"]
self._release_finished_requests(finished_requests_ids) self._release_finished_requests(finished_requests_ids)
mamba_cache_tensors = self._prepare_current_run_mamba_cache( mamba_cache_tensors = self._prepare_current_run_mamba_cache(
request_ids_to_seq_ids, finished_requests_ids) request_ids_to_seq_ids, finished_requests_ids)
else: else:
# CUDA graph capturing runs # CUDA graph capturing runs
mamba_cache_tensors = kwargs["seqlen_agnostic_capture_inputs"] mamba_cache_tensors = kwargs["seqlen_agnostic_capture_inputs"]
return mamba_cache_tensors return mamba_cache_tensors
def copy_inputs_before_cuda_graphs(self, input_buffers, **kwargs): def copy_inputs_before_cuda_graphs(self, input_buffers, **kwargs):
""" """
Copy the relevant Mamba cache into the CUDA graph input buffer Copy the relevant Mamba cache into the CUDA graph input buffer
that was provided during the capture runs that was provided during the capture runs
(JambaForCausalLM.mamba_gc_cache_buffer). (JambaForCausalLM.mamba_gc_cache_buffer).
""" """
assert all( assert all(
key in kwargs key in kwargs
for key in ["request_ids_to_seq_ids", "finished_requests_ids"]) for key in ["request_ids_to_seq_ids", "finished_requests_ids"])
finished_requests_ids = kwargs["finished_requests_ids"] finished_requests_ids = kwargs["finished_requests_ids"]
request_ids_to_seq_ids = kwargs["request_ids_to_seq_ids"] request_ids_to_seq_ids = kwargs["request_ids_to_seq_ids"]
self._release_finished_requests(finished_requests_ids) self._release_finished_requests(finished_requests_ids)
self._prepare_current_run_mamba_cache(request_ids_to_seq_ids, self._prepare_current_run_mamba_cache(request_ids_to_seq_ids,
finished_requests_ids) finished_requests_ids)
def get_seqlen_agnostic_capture_inputs(self, batch_size: int): def get_seqlen_agnostic_capture_inputs(self, batch_size: int):
""" """
Provide the CUDA graph capture runs with a buffer in adjusted size. Provide the CUDA graph capture runs with a buffer in adjusted size.
The buffer is used to maintain the Mamba Cache during the CUDA graph The buffer is used to maintain the Mamba Cache during the CUDA graph
replay runs. replay runs.
""" """
return tuple(buffer[:, :batch_size] for buffer in self.mamba_cache) return tuple(buffer[:, :batch_size] for buffer in self.mamba_cache)
def _swap_mamba_cache(self, from_index: int, to_index: int): def _swap_mamba_cache(self, from_index: int, to_index: int):
assert len(self.mamba_cache) > 0 assert len(self.mamba_cache) > 0
for cache_t in self.mamba_cache: for cache_t in self.mamba_cache:
cache_t[:, [to_index,from_index]] = \ cache_t[:, [to_index,from_index]] = \
cache_t[:, [from_index,to_index]] cache_t[:, [from_index,to_index]]
def _copy_mamba_cache(self, from_index: int, to_index: int): def _copy_mamba_cache(self, from_index: int, to_index: int):
assert len(self.mamba_cache) > 0 assert len(self.mamba_cache) > 0
for cache_t in self.mamba_cache: for cache_t in self.mamba_cache:
cache_t[:, to_index].copy_(cache_t[:, from_index], cache_t[:, to_index].copy_(cache_t[:, from_index],
non_blocking=True) non_blocking=True)
def _move_out_if_already_occupied(self, index: int, def _move_out_if_already_occupied(self, index: int,
all_occupied_indices: List[int]): all_occupied_indices: List[int]):
if index in all_occupied_indices: if index in all_occupied_indices:
first_free_index = self._first_free_index_in_mamba_cache() first_free_index = self._first_free_index_in_mamba_cache()
# In case occupied, move the occupied to a new empty block # In case occupied, move the occupied to a new empty block
self._move_cache_index_and_mappings(from_index=index, self._move_cache_index_and_mappings(from_index=index,
to_index=first_free_index) to_index=first_free_index)
def _assign_seq_id_to_mamba_cache_in_specific_dest(self, cur_rid: str, def _assign_seq_id_to_mamba_cache_in_specific_dest(self, cur_rid: str,
seq_id: int, seq_id: int,
destination_index: int): destination_index: int):
""" """
Assign (req_id,seq_id) pair to a `destination_index` index, if Assign (req_id,seq_id) pair to a `destination_index` index, if
already occupied, move the occupying index to a free index. already occupied, move the occupying index to a free index.
""" """
all_occupied_indices = self._get_all_occupied_indices() all_occupied_indices = self._get_all_occupied_indices()
if cur_rid not in self.mamba_cache_indices_mapping: if cur_rid not in self.mamba_cache_indices_mapping:
self._move_out_if_already_occupied( self._move_out_if_already_occupied(
index=destination_index, index=destination_index,
all_occupied_indices=all_occupied_indices) all_occupied_indices=all_occupied_indices)
for cache_t in self.mamba_cache: for cache_t in self.mamba_cache:
cache_t[:, destination_index].zero_() cache_t[:, destination_index].zero_()
self.mamba_cache_indices_mapping[cur_rid] = { self.mamba_cache_indices_mapping[cur_rid] = {
seq_id: destination_index seq_id: destination_index
} }
elif seq_id not in (seq_ids2indices := elif seq_id not in (seq_ids2indices :=
self.mamba_cache_indices_mapping[cur_rid]): self.mamba_cache_indices_mapping[cur_rid]):
# parallel sampling , where n > 1, assume prefill have # parallel sampling , where n > 1, assume prefill have
# already happened now we only need to copy the already # already happened now we only need to copy the already
# existing cache into the siblings seq_ids caches # existing cache into the siblings seq_ids caches
self._move_out_if_already_occupied( self._move_out_if_already_occupied(
index=destination_index, index=destination_index,
all_occupied_indices=all_occupied_indices) all_occupied_indices=all_occupied_indices)
index_exists = list(seq_ids2indices.values())[0] index_exists = list(seq_ids2indices.values())[0]
# case of decoding n>1, copy prefill cache to decoding indices # case of decoding n>1, copy prefill cache to decoding indices
self._copy_mamba_cache(from_index=index_exists, self._copy_mamba_cache(from_index=index_exists,
to_index=destination_index) to_index=destination_index)
self.mamba_cache_indices_mapping[cur_rid][ self.mamba_cache_indices_mapping[cur_rid][
seq_id] = destination_index seq_id] = destination_index
else: else:
# already exists # already exists
cache_index_already_exists = self.mamba_cache_indices_mapping[ cache_index_already_exists = self.mamba_cache_indices_mapping[
cur_rid][seq_id] cur_rid][seq_id]
if cache_index_already_exists != destination_index: if cache_index_already_exists != destination_index:
# In case the seq id already exists but not in # In case the seq id already exists but not in
# the right destination, swap it with what's occupying it # the right destination, swap it with what's occupying it
self._swap_pair_indices_and_mappings( self._swap_pair_indices_and_mappings(
from_index=cache_index_already_exists, from_index=cache_index_already_exists,
to_index=destination_index) to_index=destination_index)
def _prepare_current_run_mamba_cache( def _prepare_current_run_mamba_cache(
self, request_ids_to_seq_ids: Dict[str, list[int]], self, request_ids_to_seq_ids: Dict[str, list[int]],
finished_requests_ids: List[str]): finished_requests_ids: List[str]):
running_indices = [] running_indices = []
request_ids_to_seq_ids_flatten = [ request_ids_to_seq_ids_flatten = [
(req_id, seq_id) (req_id, seq_id)
for req_id, seq_ids in request_ids_to_seq_ids.items() for req_id, seq_ids in request_ids_to_seq_ids.items()
for seq_id in seq_ids for seq_id in seq_ids
] ]
batch_size = len(request_ids_to_seq_ids_flatten) batch_size = len(request_ids_to_seq_ids_flatten)
for dest_index, (request_id, for dest_index, (request_id,
seq_id) in enumerate(request_ids_to_seq_ids_flatten): seq_id) in enumerate(request_ids_to_seq_ids_flatten):
if request_id in finished_requests_ids: if request_id in finished_requests_ids:
# Do not allocate cache index for requests that run # Do not allocate cache index for requests that run
# and finish right after # and finish right after
continue continue
self._assign_seq_id_to_mamba_cache_in_specific_dest( self._assign_seq_id_to_mamba_cache_in_specific_dest(
request_id, seq_id, dest_index) request_id, seq_id, dest_index)
running_indices.append(dest_index) running_indices.append(dest_index)
self._clean_up_first_bs_blocks(batch_size, running_indices) self._clean_up_first_bs_blocks(batch_size, running_indices)
conv_state = self.mamba_cache[0][:, :batch_size] conv_state = self.mamba_cache[0][:, :batch_size]
temporal_state = self.mamba_cache[1][:, :batch_size] temporal_state = self.mamba_cache[1][:, :batch_size]
return (conv_state, temporal_state) return (conv_state, temporal_state)
def _get_all_occupied_indices(self): def _get_all_occupied_indices(self):
return [ return [
cache_idx cache_idx
for seq_ids2indices in self.mamba_cache_indices_mapping.values() for seq_ids2indices in self.mamba_cache_indices_mapping.values()
for cache_idx in seq_ids2indices.values() for cache_idx in seq_ids2indices.values()
] ]
def _clean_up_first_bs_blocks(self, batch_size: int, def _clean_up_first_bs_blocks(self, batch_size: int,
indices_for_current_run: List[int]): indices_for_current_run: List[int]):
# move out all of the occupied but currently not running blocks # move out all of the occupied but currently not running blocks
# outside of the first n blocks # outside of the first n blocks
destination_indices = range(batch_size) destination_indices = range(batch_size)
max_possible_batch_size = self.mamba_cache[0].shape[1] max_possible_batch_size = self.mamba_cache[0].shape[1]
for destination_index in destination_indices: for destination_index in destination_indices:
if destination_index in self._get_all_occupied_indices() and \ if destination_index in self._get_all_occupied_indices() and \
destination_index not in indices_for_current_run: destination_index not in indices_for_current_run:
# move not running indices outside of the batch # move not running indices outside of the batch
all_other_indices = list( all_other_indices = list(
range(batch_size, max_possible_batch_size)) range(batch_size, max_possible_batch_size))
first_avail_index = self._first_free_index_in_mamba_cache( first_avail_index = self._first_free_index_in_mamba_cache(
all_other_indices) all_other_indices)
self._swap_indices(from_index=destination_index, self._swap_indices(from_index=destination_index,
to_index=first_avail_index) to_index=first_avail_index)
def _move_cache_index_and_mappings(self, from_index: int, to_index: int): def _move_cache_index_and_mappings(self, from_index: int, to_index: int):
self._copy_mamba_cache(from_index=from_index, to_index=to_index) self._copy_mamba_cache(from_index=from_index, to_index=to_index)
self._update_mapping_index(from_index=from_index, to_index=to_index) self._update_mapping_index(from_index=from_index, to_index=to_index)
def _swap_pair_indices_and_mappings(self, from_index: int, to_index: int): def _swap_pair_indices_and_mappings(self, from_index: int, to_index: int):
self._swap_mamba_cache(from_index=from_index, to_index=to_index) self._swap_mamba_cache(from_index=from_index, to_index=to_index)
self._swap_mapping_index(from_index=from_index, to_index=to_index) self._swap_mapping_index(from_index=from_index, to_index=to_index)
def _swap_mapping_index(self, from_index: int, to_index: int): def _swap_mapping_index(self, from_index: int, to_index: int):
for seq_ids2index in self.mamba_cache_indices_mapping.values(): for seq_ids2index in self.mamba_cache_indices_mapping.values():
for seq_id, index in seq_ids2index.items(): for seq_id, index in seq_ids2index.items():
if from_index == index: if from_index == index:
seq_ids2index.update({seq_id: to_index}) seq_ids2index.update({seq_id: to_index})
elif to_index == index: elif to_index == index:
seq_ids2index.update({seq_id: from_index}) seq_ids2index.update({seq_id: from_index})
def _update_mapping_index(self, from_index: int, to_index: int): def _update_mapping_index(self, from_index: int, to_index: int):
for seq_ids2index in self.mamba_cache_indices_mapping.values(): for seq_ids2index in self.mamba_cache_indices_mapping.values():
for seq_id, index in seq_ids2index.items(): for seq_id, index in seq_ids2index.items():
if from_index == index: if from_index == index:
seq_ids2index.update({seq_id: to_index}) seq_ids2index.update({seq_id: to_index})
return return
def _release_finished_requests(self, def _release_finished_requests(self,
finished_seq_groups_req_ids: List[str]): finished_seq_groups_req_ids: List[str]):
for req_id in finished_seq_groups_req_ids: for req_id in finished_seq_groups_req_ids:
if req_id in self.mamba_cache_indices_mapping: if req_id in self.mamba_cache_indices_mapping:
self.mamba_cache_indices_mapping.pop(req_id) self.mamba_cache_indices_mapping.pop(req_id)
def _first_free_index_in_mamba_cache( def _first_free_index_in_mamba_cache(
self, indices_range: Optional[List[int]] = None) -> int: self, indices_range: Optional[List[int]] = None) -> int:
assert self.mamba_cache is not None assert self.mamba_cache is not None
if indices_range is None: if indices_range is None:
max_possible_batch_size = self.mamba_cache[0].shape[1] max_possible_batch_size = self.mamba_cache[0].shape[1]
indices_range = list(range(max_possible_batch_size)) indices_range = list(range(max_possible_batch_size))
all_occupied_indices = self._get_all_occupied_indices() all_occupied_indices = self._get_all_occupied_indices()
for i in indices_range: for i in indices_range:
if i not in all_occupied_indices: if i not in all_occupied_indices:
return i return i
raise Exception("Couldn't find a free spot in the mamba cache! This" raise Exception("Couldn't find a free spot in the mamba cache! This"
"should never happen") "should never happen")

View File

@@ -1,11 +1,11 @@
""" """
Patches the vLLM model registry and deploys the Qwen3_5 model file. Patches the vLLM model registry and deploys the Qwen3_5 model file.
Deploy steps on the remote machine: Deploy steps on the remote machine:
1. patch_ops.sh locates vLLM with importlib.util.find_spec. 1. patch_ops.sh locates vLLM with importlib.util.find_spec.
2. cp modified_scripts/qwen3_5.py into the detected vllm model directory. 2. cp modified_scripts/qwen3_5.py into the detected vllm model directory.
2. python3 modified_scripts/patch_vllm_qwen3_5.py 2. python3 modified_scripts/patch_vllm_qwen3_5.py
The registry patch installs Qwen3.6 aliases so /model/config.json does not The registry patch installs Qwen3.6 aliases so /model/config.json does not
need to be edited by hand. need to be edited by hand.
""" """
@@ -26,8 +26,8 @@ EXPECTED_REGISTRY_ENTRIES = (
'"Qwen3_6ForCausalLM": ("qwen3_5", "Qwen3_5ForCausalLM")', '"Qwen3_6ForCausalLM": ("qwen3_5", "Qwen3_5ForCausalLM")',
'"Qwen3_6MoeForCausalLM": ("qwen3_5", "Qwen3_5MoeForCausalLM")', '"Qwen3_6MoeForCausalLM": ("qwen3_5", "Qwen3_5MoeForCausalLM")',
) )
def main(): def main():
print(f"=== Patching {REGISTRY} ===") print(f"=== Patching {REGISTRY} ===")
replace_once( replace_once(
@@ -42,7 +42,7 @@ def main():
' "Qwen3_6MoeForCausalLM": ("qwen3_5", "Qwen3_5MoeForCausalLM"),', ' "Qwen3_6MoeForCausalLM": ("qwen3_5", "Qwen3_5MoeForCausalLM"),',
required=True, required=True,
already_contains='"Qwen3_6MoeForCausalLM"') already_contains='"Qwen3_6MoeForCausalLM"')
print("\n=== Static verification ===") print("\n=== Static verification ===")
model_source = MODEL.read_text(encoding="utf-8") model_source = MODEL.read_text(encoding="utf-8")
tree = ast.parse(model_source, filename=str(MODEL)) tree = ast.parse(model_source, filename=str(MODEL))
@@ -67,7 +67,7 @@ def main():
print(f" registry aliases verified: {len(EXPECTED_REGISTRY_ENTRIES)}") print(f" registry aliases verified: {len(EXPECTED_REGISTRY_ENTRIES)}")
print("\nDone. Registry aliases installed; do not edit /model/config.json.") print("\nDone. Registry aliases installed; do not edit /model/config.json.")
if __name__ == "__main__": if __name__ == "__main__":
main() main()

View File

@@ -1,23 +1,23 @@
""" """
Patches vLLM 0.6.3 to register Qwen3CoderToolParser under the name "qwen3_coder". Patches vLLM 0.6.3 to register Qwen3CoderToolParser under the name "qwen3_coder".
Deploy steps on the remote machine (already called by patch_ops.sh): Deploy steps on the remote machine (already called by patch_ops.sh):
1. patch_ops.sh locates vLLM with importlib.util.find_spec. 1. patch_ops.sh locates vLLM with importlib.util.find_spec.
2. cp qwen3coder_tool_parser.py into the detected vllm tool_parsers. 2. cp qwen3coder_tool_parser.py into the detected vllm tool_parsers.
2. python3 patch_vllm_tool_parser.py 2. python3 patch_vllm_tool_parser.py
Usage after patching: Usage after patching:
--tool-call-parser qwen3_coder --enable-auto-tool-choice --tool-call-parser qwen3_coder --enable-auto-tool-choice
""" """
from patch_utils import ensure_dir, package_root, replace_once from patch_utils import ensure_dir, package_root, replace_once
VLLM_ROOT = package_root("vllm") VLLM_ROOT = package_root("vllm")
TOOL_PARSERS_DIR = VLLM_ROOT / "entrypoints" / "openai" / "tool_parsers" TOOL_PARSERS_DIR = VLLM_ROOT / "entrypoints" / "openai" / "tool_parsers"
INIT_FILE = TOOL_PARSERS_DIR / "__init__.py" INIT_FILE = TOOL_PARSERS_DIR / "__init__.py"
def main(): def main():
ensure_dir(TOOL_PARSERS_DIR) ensure_dir(TOOL_PARSERS_DIR)
print(f"=== Patching {INIT_FILE} ===") print(f"=== Patching {INIT_FILE} ===")
@@ -35,23 +35,23 @@ def main():
' "Qwen3CoderToolParser"\n]', ' "Qwen3CoderToolParser"\n]',
required=True, required=True,
already_contains='"Qwen3CoderToolParser"') already_contains='"Qwen3CoderToolParser"')
print("\n=== Verification ===") print("\n=== Verification ===")
try: try:
import importlib.util import importlib.util
spec = importlib.util.spec_from_file_location( spec = importlib.util.spec_from_file_location(
"qwen3coder_tool_parser", "qwen3coder_tool_parser",
str(TOOL_PARSERS_DIR / "qwen3coder_tool_parser.py"), str(TOOL_PARSERS_DIR / "qwen3coder_tool_parser.py"),
) )
mod = importlib.util.module_from_spec(spec) mod = importlib.util.module_from_spec(spec)
print(f" Module spec loaded: {spec.name}") print(f" Module spec loaded: {spec.name}")
print(" (full import requires torch/vllm runtime — skipping exec)") print(" (full import requires torch/vllm runtime — skipping exec)")
except Exception as e: except Exception as e:
print(f" [optional] spec check failed: {e}") print(f" [optional] spec check failed: {e}")
print("\nDone. Start vLLM server with:") print("\nDone. Start vLLM server with:")
print(" --tool-call-parser qwen3_coder --enable-auto-tool-choice") print(" --tool-call-parser qwen3_coder --enable-auto-tool-choice")
if __name__ == "__main__": if __name__ == "__main__":
main() main()

View File

@@ -1,29 +1,29 @@
""" """
策略:批量(block-diagonal)fallback — 纯 PyTorch 数学实现 策略:批量(block-diagonal)fallback — 纯 PyTorch 数学实现
============================================================= =============================================================
构建块对角 causal mask,对整批序列一次 matmul + softmax, 构建块对角 causal mask,对整批序列一次 matmul + softmax,
完全绕开所有硬件 flash attention kernel。 完全绕开所有硬件 flash attention kernel。
背景: 背景:
ixformer flshattF: head_dim > 128 报错拒绝 ixformer flshattF: head_dim > 128 报错拒绝
cudnnFlashAttnForward: 接受 head_dim=256,但数值结果错误(输出全"!") cudnnFlashAttnForward: 接受 head_dim=256,但数值结果错误(输出全"!")
两者大概率是同一硬件单元,ixformer 提前拦截了硬件不支持的配置。 两者大概率是同一硬件单元,ixformer 提前拦截了硬件不支持的配置。
纯 matmul 路径完全绕开硬件 flash attention,数值正确。 纯 matmul 路径完全绕开硬件 flash attention,数值正确。
优点: 优点:
数值正确。 数值正确。
并发请求 prefill attention 在 GPU 上真正并行(一次大 matmul)。 并发请求 prefill attention 在 GPU 上真正并行(一次大 matmul)。
缺点: 缺点:
峰值显存 = total_tokens² × H × dtype_size 峰值显存 = total_tokens² × H × dtype_size
total_tokens 受 --max-num-batched-tokens 控制,max-model-len 控制不住。 total_tokens 受 --max-num-batched-tokens 控制,max-model-len 控制不住。
内存参考(fp16,H_local=6,--max-num-batched-tokens=T): 内存参考(fp16,H_local=6,--max-num-batched-tokens=T):
T=2048 → 峰值 ~50 MB T=2048 → 峰值 ~50 MB
T=4096 → 峰值 ~200 MB T=4096 → 峰值 ~200 MB
T=8192 → 峰值 ~800 MB T=8192 → 峰值 ~800 MB
T=16384 → 峰值 ~3.2 GB T=16384 → 峰值 ~3.2 GB
Deploy: Deploy:
python3 modified_scripts/patch_xformers_sdpa_batch.py python3 modified_scripts/patch_xformers_sdpa_batch.py
""" """
@@ -31,126 +31,126 @@ Deploy:
from patch_utils import package_root, replace_once from patch_utils import package_root, replace_once
XFORMERS_PATH = package_root("vllm") / "attention" / "backends" / "xformers.py" XFORMERS_PATH = package_root("vllm") / "attention" / "backends" / "xformers.py"
FALLBACK_METHOD = ''' FALLBACK_METHOD = '''
def _run_sdpa_fallback( def _run_sdpa_fallback(
self, self,
query: torch.Tensor, query: torch.Tensor,
key: torch.Tensor, key: torch.Tensor,
value: torch.Tensor, value: torch.Tensor,
attn_metadata: "XFormersMetadata", attn_metadata: "XFormersMetadata",
) -> torch.Tensor: ) -> torch.Tensor:
"""批量纯数学 attention fallback。 """批量纯数学 attention fallback。
构建块对角 causal mask(等价于 ixformer BlockDiagonalCausalMask), 构建块对角 causal mask(等价于 ixformer BlockDiagonalCausalMask),
对整批序列一次 matmul + softmax,GPU 并行处理所有序列。 对整批序列一次 matmul + softmax,GPU 并行处理所有序列。
块对角 mask 结构(seq1 len=3,seq2 len=2): 块对角 mask 结构(seq1 len=3,seq2 len=2):
s1,0 s1,1 s1,2 s2,0 s2,1 s1,0 s1,1 s1,2 s2,0 s2,1
s1,0 [ 0 -inf -inf -inf -inf ] s1,0 [ 0 -inf -inf -inf -inf ]
s1,1 [ 0 0 -inf -inf -inf ] s1,1 [ 0 0 -inf -inf -inf ]
s1,2 [ 0 0 0 -inf -inf ] s1,2 [ 0 0 0 -inf -inf ]
s2,0 [-inf -inf -inf 0 -inf ] s2,0 [-inf -inf -inf 0 -inf ]
s2,1 [-inf -inf -inf 0 0 ] s2,1 [-inf -inf -inf 0 0 ]
softmax 在 float32 下计算防止 float16 溢出,结果转回原始 dtype。 softmax 在 float32 下计算防止 float16 溢出,结果转回原始 dtype。
Args: Args:
query : [1, total_prefill_tokens, num_heads, head_dim] query : [1, total_prefill_tokens, num_heads, head_dim]
key : [1, total_prefill_tokens, num_kv_heads, head_dim] key : [1, total_prefill_tokens, num_kv_heads, head_dim]
value : [1, total_prefill_tokens, num_kv_heads, head_dim] value : [1, total_prefill_tokens, num_kv_heads, head_dim]
Returns: Returns:
[1, total_prefill_tokens, num_heads, head_dim] [1, total_prefill_tokens, num_heads, head_dim]
""" """
assert attn_metadata.seq_lens is not None assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype orig_dtype = query.dtype
total_tokens = query.shape[1] total_tokens = query.shape[1]
# ── 构建块对角 causal mask [T, T] ──────────────────────────────── # ── 构建块对角 causal mask [T, T] ────────────────────────────────
# 全部初始化为 -inf,再对每条序列的对角块填入下三角 0 # 全部初始化为 -inf,再对每条序列的对角块填入下三角 0
mask = torch.full( mask = torch.full(
(total_tokens, total_tokens), (total_tokens, total_tokens),
float("-inf"), float("-inf"),
dtype=torch.float32, dtype=torch.float32,
device=query.device, device=query.device,
) )
start = 0 start = 0
for seq_len in attn_metadata.seq_lens: for seq_len in attn_metadata.seq_lens:
end = start + seq_len end = start + seq_len
mask[start:end, start:end] = torch.tril( mask[start:end, start:end] = torch.tril(
torch.zeros(seq_len, seq_len, torch.zeros(seq_len, seq_len,
dtype=torch.float32, device=query.device) dtype=torch.float32, device=query.device)
) )
start = end start = end
# ── [1, H, T, D],.contiguous() ────────────────────────────────── # ── [1, H, T, D],.contiguous() ──────────────────────────────────
q_all = query.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0) q_all = query.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
k_all = key.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0) k_all = key.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
v_all = value.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0) v_all = value.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
# ── GQA:展开 KV heads ──────────────────────────────────────────── # ── GQA:展开 KV heads ────────────────────────────────────────────
if k_all.shape[1] != q_all.shape[1]: if k_all.shape[1] != q_all.shape[1]:
n = q_all.shape[1] // k_all.shape[1] n = q_all.shape[1] // k_all.shape[1]
k_all = k_all.repeat_interleave(n, dim=1).contiguous() k_all = k_all.repeat_interleave(n, dim=1).contiguous()
v_all = v_all.repeat_interleave(n, dim=1).contiguous() v_all = v_all.repeat_interleave(n, dim=1).contiguous()
# ── 纯数学 attention(float32 防溢出)──────────────────────────── # ── 纯数学 attention(float32 防溢出)────────────────────────────
# [1, H, T, T] # [1, H, T, T]
attn_w = torch.matmul(q_all.float(), k_all.float().transpose(-2, -1)) attn_w = torch.matmul(q_all.float(), k_all.float().transpose(-2, -1))
attn_w = attn_w * self.scale attn_w = attn_w * self.scale
attn_w = attn_w + mask # 加法广播:mask [T,T] → [1, H, T, T] attn_w = attn_w + mask # 加法广播:mask [T,T] → [1, H, T, T]
attn_w = torch.softmax(attn_w, dim=-1) attn_w = torch.softmax(attn_w, dim=-1)
out = torch.matmul(attn_w, v_all.float()).to(orig_dtype) out = torch.matmul(attn_w, v_all.float()).to(orig_dtype)
# [1, H, T, D] → [1, T, H, D] # [1, H, T, D] → [1, T, H, D]
return out.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0) return out.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
''' '''
OLD_XFORMER_BLOCK = """\ OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp() self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None: if self.alibi_slopes is None:
# Add the batch dimension. # Add the batch dimension.
query = query.unsqueeze(0) query = query.unsqueeze(0)
key = key.unsqueeze(0) key = key.unsqueeze(0)
value = value.unsqueeze(0) value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward( out = xops.memory_efficient_attention_forward(
query, query,
key, key,
value, value,
attn_bias=attn_bias[0], attn_bias=attn_bias[0],
p=0.0, p=0.0,
scale=self.scale, scale=self.scale,
op = self.attn_op op = self.attn_op
) )
return out.view_as(original_query)\ return out.view_as(original_query)\
""" """
NEW_XFORMER_BLOCK = """\ NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp() self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None: if self.alibi_slopes is None:
# Add the batch dimension. # Add the batch dimension.
query = query.unsqueeze(0) query = query.unsqueeze(0)
key = key.unsqueeze(0) key = key.unsqueeze(0)
value = value.unsqueeze(0) value = value.unsqueeze(0)
if self.head_size > 128: if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata) out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else: else:
out = xops.memory_efficient_attention_forward( out = xops.memory_efficient_attention_forward(
query, query,
key, key,
value, value,
attn_bias=attn_bias[0], attn_bias=attn_bias[0],
p=0.0, p=0.0,
scale=self.scale, scale=self.scale,
op=self.attn_op, op=self.attn_op,
) )
return out.view_as(original_query)\ return out.view_as(original_query)\
""" """
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward(" INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
def patch_file(path): def patch_file(path):
replace_once( replace_once(
path, path,
@@ -164,14 +164,14 @@ def patch_file(path):
NEW_XFORMER_BLOCK, NEW_XFORMER_BLOCK,
required=True, required=True,
already_contains="out = self._run_sdpa_fallback(query, key, value, attn_metadata)") already_contains="out = self._run_sdpa_fallback(query, key, value, attn_metadata)")
def main(): def main():
print("=== patch_xformers_sdpa_batch (batch, pure-math) ===") print("=== patch_xformers_sdpa_batch (batch, pure-math) ===")
print(f"Target: {XFORMERS_PATH}") print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH) patch_file(XFORMERS_PATH)
print("\nDone.") print("\nDone.")
if __name__ == "__main__": if __name__ == "__main__":
main() main()

View File

@@ -1,26 +1,26 @@
""" """
策略:批量(block-diagonal)— F.scaled_dot_product_attention,可走硬件 kernel 策略:批量(block-diagonal)— F.scaled_dot_product_attention,可走硬件 kernel
============================================================================= =============================================================================
构建块对角 causal mask,对整批序列一次 F.scaled_dot_product_attention。 构建块对角 causal mask,对整批序列一次 F.scaled_dot_product_attention。
与 patch_xformers_sdpa_batch.py(纯 matmul)的区别: 与 patch_xformers_sdpa_batch.py(纯 matmul)的区别:
SDPA 会根据 PyTorch/驱动能力分发到最优 kernel(Flash Attention / SDPA 会根据 PyTorch/驱动能力分发到最优 kernel(Flash Attention /
mem-efficient attention / math fallback),而不是固定走 cublas matmul。 mem-efficient attention / math fallback),而不是固定走 cublas matmul。
历史说明: 历史说明:
该方案最早因输出全"!"而被弃用,后续排查确认"!"由 mamba_cache.py bug 该方案最早因输出全"!"而被弃用,后续排查确认"!"由 mamba_cache.py bug
引起,与 attention 实现无关。当前恢复此方案用于性能对比测试。 引起,与 attention 实现无关。当前恢复此方案用于性能对比测试。
已知硬件限制(BI-V100): 已知硬件限制(BI-V100):
cudnnFlashAttnForward 不支持 is_causal=True(报错)。 cudnnFlashAttnForward 不支持 is_causal=True(报错)。
本实现使用 is_causal=False + 显式块对角 additive mask 规避此限制。 本实现使用 is_causal=False + 显式块对角 additive mask 规避此限制。
若 SDPA 仍分发到有问题的 kernel,回退到 patch_xformers_sdpa_batch.py。 若 SDPA 仍分发到有问题的 kernel,回退到 patch_xformers_sdpa_batch.py。
优点(vs 纯 matmul): 优点(vs 纯 matmul):
SDPA 可分发到 Flash Attention kernel → O(L) 显存、更快的 CUDA kernel。 SDPA 可分发到 Flash Attention kernel → O(L) 显存、更快的 CUDA kernel。
缺点: 缺点:
依赖硬件 kernel 行为,若 kernel 有 bug 则数值错误(需与 matmul 版对比验证)。 依赖硬件 kernel 行为,若 kernel 有 bug 则数值错误(需与 matmul 版对比验证)。
Deploy: Deploy:
python3 modified_scripts/patch_xformers_sdpa_batch_kernel.py python3 modified_scripts/patch_xformers_sdpa_batch_kernel.py
""" """
@@ -28,128 +28,128 @@ Deploy:
from patch_utils import package_root, replace_once from patch_utils import package_root, replace_once
XFORMERS_PATH = package_root("vllm") / "attention" / "backends" / "xformers.py" XFORMERS_PATH = package_root("vllm") / "attention" / "backends" / "xformers.py"
FALLBACK_METHOD = ''' FALLBACK_METHOD = '''
def _run_sdpa_fallback( def _run_sdpa_fallback(
self, self,
query: torch.Tensor, query: torch.Tensor,
key: torch.Tensor, key: torch.Tensor,
value: torch.Tensor, value: torch.Tensor,
attn_metadata: "XFormersMetadata", attn_metadata: "XFormersMetadata",
) -> torch.Tensor: ) -> torch.Tensor:
"""批量 F.scaled_dot_product_attention fallback(可走硬件 kernel)。 """批量 F.scaled_dot_product_attention fallback(可走硬件 kernel)。
构建块对角 causal mask,对整批序列一次 SDPA 调用。 构建块对角 causal mask,对整批序列一次 SDPA 调用。
SDPA 可分发到 Flash Attention / mem-efficient attention kernel。 SDPA 可分发到 Flash Attention / mem-efficient attention kernel。
is_causal=False + 显式 additive mask,规避 cudnnFlashAttnForward is_causal=False + 显式 additive mask,规避 cudnnFlashAttnForward
不支持 is_causal=True 的限制。 不支持 is_causal=True 的限制。
块对角 mask(seq1 len=3,seq2 len=2): 块对角 mask(seq1 len=3,seq2 len=2):
s1,0 s1,1 s1,2 s2,0 s2,1 s1,0 s1,1 s1,2 s2,0 s2,1
s1,0 [ 0 -inf -inf -inf -inf ] s1,0 [ 0 -inf -inf -inf -inf ]
s1,1 [ 0 0 -inf -inf -inf ] s1,1 [ 0 0 -inf -inf -inf ]
s1,2 [ 0 0 0 -inf -inf ] s1,2 [ 0 0 0 -inf -inf ]
s2,0 [-inf -inf -inf 0 -inf ] s2,0 [-inf -inf -inf 0 -inf ]
s2,1 [-inf -inf -inf 0 0 ] s2,1 [-inf -inf -inf 0 0 ]
Args: Args:
query : [1, total_prefill_tokens, num_heads, head_dim] query : [1, total_prefill_tokens, num_heads, head_dim]
key : [1, total_prefill_tokens, num_kv_heads, head_dim] key : [1, total_prefill_tokens, num_kv_heads, head_dim]
value : [1, total_prefill_tokens, num_kv_heads, head_dim] value : [1, total_prefill_tokens, num_kv_heads, head_dim]
Returns: Returns:
[1, total_prefill_tokens, num_heads, head_dim] [1, total_prefill_tokens, num_heads, head_dim]
""" """
import torch.nn.functional as F import torch.nn.functional as F
assert attn_metadata.seq_lens is not None assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype orig_dtype = query.dtype
total_tokens = query.shape[1] total_tokens = query.shape[1]
# ── 块对角 causal mask [T, T] ───────────────────────────────────── # ── 块对角 causal mask [T, T] ─────────────────────────────────────
mask = torch.full( mask = torch.full(
(total_tokens, total_tokens), (total_tokens, total_tokens),
float("-inf"), float("-inf"),
dtype=orig_dtype, dtype=orig_dtype,
device=query.device, device=query.device,
) )
start = 0 start = 0
for seq_len in attn_metadata.seq_lens: for seq_len in attn_metadata.seq_lens:
end = start + seq_len end = start + seq_len
mask[start:end, start:end] = torch.tril( mask[start:end, start:end] = torch.tril(
torch.zeros(seq_len, seq_len, dtype=orig_dtype, device=query.device) torch.zeros(seq_len, seq_len, dtype=orig_dtype, device=query.device)
) )
start = end start = end
# ── [1, H, T, D] ────────────────────────────────────────────────── # ── [1, H, T, D] ──────────────────────────────────────────────────
q_all = query.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0) q_all = query.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
k_all = key.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0) k_all = key.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
v_all = value.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0) v_all = value.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
# ── GQA:展开 KV heads ──────────────────────────────────────────── # ── GQA:展开 KV heads ────────────────────────────────────────────
if k_all.shape[1] != q_all.shape[1]: if k_all.shape[1] != q_all.shape[1]:
n = q_all.shape[1] // k_all.shape[1] n = q_all.shape[1] // k_all.shape[1]
k_all = k_all.repeat_interleave(n, dim=1).contiguous() k_all = k_all.repeat_interleave(n, dim=1).contiguous()
v_all = v_all.repeat_interleave(n, dim=1).contiguous() v_all = v_all.repeat_interleave(n, dim=1).contiguous()
# ── F.scaled_dot_product_attention(可走硬件 kernel)───────────── # ── F.scaled_dot_product_attention(可走硬件 kernel)─────────────
# is_causal=False:避免 cudnnFlashAttnForward "not support causal mode" # is_causal=False:避免 cudnnFlashAttnForward "not support causal mode"
# attn_mask 传 additive float mask(非 bool),SDPA 选择 math/kernel 路径 # attn_mask 传 additive float mask(非 bool),SDPA 选择 math/kernel 路径
out = F.scaled_dot_product_attention( out = F.scaled_dot_product_attention(
q_all, k_all, v_all, q_all, k_all, v_all,
attn_mask=mask, attn_mask=mask,
dropout_p=0.0, dropout_p=0.0,
is_causal=False, is_causal=False,
scale=self.scale, scale=self.scale,
) )
# [1, H, T, D] → [1, T, H, D] # [1, H, T, D] → [1, T, H, D]
return out.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0) return out.squeeze(0).permute(1, 0, 2).contiguous().unsqueeze(0)
''' '''
OLD_XFORMER_BLOCK = """\ OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp() self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None: if self.alibi_slopes is None:
# Add the batch dimension. # Add the batch dimension.
query = query.unsqueeze(0) query = query.unsqueeze(0)
key = key.unsqueeze(0) key = key.unsqueeze(0)
value = value.unsqueeze(0) value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward( out = xops.memory_efficient_attention_forward(
query, query,
key, key,
value, value,
attn_bias=attn_bias[0], attn_bias=attn_bias[0],
p=0.0, p=0.0,
scale=self.scale, scale=self.scale,
op = self.attn_op op = self.attn_op
) )
return out.view_as(original_query)\ return out.view_as(original_query)\
""" """
NEW_XFORMER_BLOCK = """\ NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp() self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None: if self.alibi_slopes is None:
# Add the batch dimension. # Add the batch dimension.
query = query.unsqueeze(0) query = query.unsqueeze(0)
key = key.unsqueeze(0) key = key.unsqueeze(0)
value = value.unsqueeze(0) value = value.unsqueeze(0)
if self.head_size > 128: if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata) out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else: else:
out = xops.memory_efficient_attention_forward( out = xops.memory_efficient_attention_forward(
query, query,
key, key,
value, value,
attn_bias=attn_bias[0], attn_bias=attn_bias[0],
p=0.0, p=0.0,
scale=self.scale, scale=self.scale,
op=self.attn_op, op=self.attn_op,
) )
return out.view_as(original_query)\ return out.view_as(original_query)\
""" """
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward(" INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
def patch_file(path): def patch_file(path):
replace_once( replace_once(
path, path,
@@ -163,14 +163,14 @@ def patch_file(path):
NEW_XFORMER_BLOCK, NEW_XFORMER_BLOCK,
required=True, required=True,
already_contains="out = self._run_sdpa_fallback(query, key, value, attn_metadata)") already_contains="out = self._run_sdpa_fallback(query, key, value, attn_metadata)")
def main(): def main():
print("=== patch_xformers_sdpa_batch_kernel (batch, F.sdpa + kernel dispatch) ===") print("=== patch_xformers_sdpa_batch_kernel (batch, F.sdpa + kernel dispatch) ===")
print(f"Target: {XFORMERS_PATH}") print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH) patch_file(XFORMERS_PATH)
print("\nDone.") print("\nDone.")
if __name__ == "__main__": if __name__ == "__main__":
main() main()

View File

@@ -1,39 +1,39 @@
""" """
策略:顺序(per-sequence)fallback — 纯 PyTorch 数学实现 策略:顺序(per-sequence)fallback — 纯 PyTorch 数学实现
========================================================== ==========================================================
逐条序列用 matmul + softmax 手写 attention,完全绕开所有硬件 逐条序列用 matmul + softmax 手写 attention,完全绕开所有硬件
flash attention kernel(ixformer / cudnnFlashAttnForward)。 flash attention kernel(ixformer / cudnnFlashAttnForward)。
背景: 背景:
Iluvatar cudnnFlashAttnForward 存在两个已知问题: Iluvatar cudnnFlashAttnForward 存在两个已知问题:
1. 不支持 is_causal=True(报错) 1. 不支持 is_causal=True(报错)
2. 使用 attn_mask 路径时数值结果不正确(静默错误,输出全为"!") 2. 使用 attn_mask 路径时数值结果不正确(静默错误,输出全为"!")
与华为昇腾 910B4 上 llama.cpp --flash-attn off 修复同类问题的原理相同。 与华为昇腾 910B4 上 llama.cpp --flash-attn off 修复同类问题的原理相同。
纯数学路径(matmul + softmax)在任何 PyTorch 后端上结果都正确。 纯数学路径(matmul + softmax)在任何 PyTorch 后端上结果都正确。
优点: 优点:
数值正确,不依赖任何硬件特定 attention kernel。 数值正确,不依赖任何硬件特定 attention kernel。
峰值显存 = max(seq_len)² × H × dtype_size,由 --max-model-len 控制。 峰值显存 = max(seq_len)² × H × dtype_size,由 --max-model-len 控制。
缺点: 缺点:
并发请求的 prefill attention 串行执行。 并发请求的 prefill attention 串行执行。
O(L²) 显存(无 flash attention 的 O(L) 优化)。 O(L²) 显存(无 flash attention 的 O(L) 优化)。
内存参考(fp16,H_local=6): 内存参考(fp16,H_local=6):
max-model-len=4096 → 峰值 ~200 MB max-model-len=4096 → 峰值 ~200 MB
max-model-len=8192 → 峰值 ~800 MB max-model-len=8192 → 峰值 ~800 MB
max-model-len=16384 → 峰值 ~3.2 GB max-model-len=16384 → 峰值 ~3.2 GB
额外 patch(arg_utils.py): 额外 patch(arg_utils.py):
vllm 0.6.3 在 max_model_len > 32K 时会自动开启 chunked prefill(无命令行 vllm 0.6.3 在 max_model_len > 32K 时会自动开启 chunked prefill(无命令行
关闭选项),原意是防止 profiling OOM。但 _run_sdpa_fallback 已通过 Q-tiling 关闭选项),原意是防止 profiling OOM。但 _run_sdpa_fallback 已通过 Q-tiling
解决了该问题,chunked prefill 反而会把推理路径从 _run_sdpa_fallback 切换到 解决了该问题,chunked prefill 反而会把推理路径从 _run_sdpa_fallback 切换到
_forward_prefix_pytorch,属于不必要的行为变更,因此一并禁用该自动逻辑。 _forward_prefix_pytorch,属于不必要的行为变更,因此一并禁用该自动逻辑。
Deploy: Deploy:
python3 modified_scripts/patch_xformers_sdpa_seq.py python3 modified_scripts/patch_xformers_sdpa_seq.py
""" """
from patch_utils import package_root, replace_one_of, replace_once from patch_utils import package_root, replace_one_of, replace_once
VLLM_ROOT = package_root("vllm") VLLM_ROOT = package_root("vllm")
@@ -44,24 +44,24 @@ LOGITS_PROC_PATH = (
OUTLINES_DECODING_PATH = ( OUTLINES_DECODING_PATH = (
VLLM_ROOT / "model_executor" / "guided_decoding" / VLLM_ROOT / "model_executor" / "guided_decoding" /
"outlines_decoding.py") "outlines_decoding.py")
# _apply_logits_processors crashes when seq_groups is None (intermediate # _apply_logits_processors crashes when seq_groups is None (intermediate
# chunked-prefill chunks on the driver rank). Add an early-return guard. # chunked-prefill chunks on the driver rank). Add an early-return guard.
_LP_OLD_BLOCK = """\ _LP_OLD_BLOCK = """\
def _apply_logits_processors( def _apply_logits_processors(
logits: torch.Tensor, logits: torch.Tensor,
sampling_metadata: SamplingMetadata, sampling_metadata: SamplingMetadata,
) -> torch.Tensor: ) -> torch.Tensor:
found_logits_processors = False\ found_logits_processors = False\
""" """
_LP_NEW_BLOCK = """\ _LP_NEW_BLOCK = """\
def _apply_logits_processors( def _apply_logits_processors(
logits: torch.Tensor, logits: torch.Tensor,
sampling_metadata: SamplingMetadata, sampling_metadata: SamplingMetadata,
) -> torch.Tensor: ) -> torch.Tensor:
if sampling_metadata.seq_groups is None: # intermediate chunked-prefill chunk if sampling_metadata.seq_groups is None: # intermediate chunked-prefill chunk
return logits return logits
found_logits_processors = False\ found_logits_processors = False\
""" """
@@ -116,26 +116,26 @@ _ws : JSON_WS?
JSON_STRING: /"(\\["\\\/bfnrt]|\\u[0-9a-fA-F]{4}|[^"\\\x00-\x1f])*"/ JSON_STRING: /"(\\["\\\/bfnrt]|\\u[0-9a-fA-F]{4}|[^"\\\x00-\x1f])*"/
JSON_WS: /[ \t\r\n]{1,4}/ JSON_WS: /[ \t\r\n]{1,4}/
%import common.SIGNED_NUMBER''' %import common.SIGNED_NUMBER'''
# vllm 0.6.3 自动开启 chunked prefill 的原始块 # vllm 0.6.3 自动开启 chunked prefill 的原始块
_ARG_OLD_BLOCK = """\ _ARG_OLD_BLOCK = """\
if (is_gpu and not use_sliding_window and not use_spec_decode if (is_gpu and not use_sliding_window and not use_spec_decode
and not self.enable_lora and not self.enable_lora
and not self.enable_prompt_adapter): and not self.enable_prompt_adapter):
self.enable_chunked_prefill = True self.enable_chunked_prefill = True
logger.warning( logger.warning(
"Chunked prefill is enabled by default for models with " "Chunked prefill is enabled by default for models with "
"max_model_len > 32K. Currently, chunked prefill might " "max_model_len > 32K. Currently, chunked prefill might "
"not work with some features or models. If you " "not work with some features or models. If you "
"encounter any issues, please disable chunked prefill " "encounter any issues, please disable chunked prefill "
"by setting --enable-chunked-prefill=False.")\ "by setting --enable-chunked-prefill=False.")\
""" """
_ARG_NEW_BLOCK = """\ _ARG_NEW_BLOCK = """\
if (is_gpu and not use_sliding_window and not use_spec_decode if (is_gpu and not use_sliding_window and not use_spec_decode
and not self.enable_lora and not self.enable_lora
and not self.enable_prompt_adapter): and not self.enable_prompt_adapter):
pass # skip auto-enable: Q-tiling in _run_sdpa_fallback pass # skip auto-enable: Q-tiling in _run_sdpa_fallback
# handles long-context memory without chunked prefill\ # handles long-context memory without chunked prefill\
""" """
@@ -163,146 +163,146 @@ _MM_PREFIX_NEW_BLOCK = """\
"supported for multimodal models and has been disabled.") "supported for multimodal models and has been disabled.")
self.enable_prefix_caching = False\ self.enable_prefix_caching = False\
""" """
FALLBACK_METHOD = ''' FALLBACK_METHOD = '''
def _run_sdpa_fallback( def _run_sdpa_fallback(
self, self,
query: torch.Tensor, query: torch.Tensor,
key: torch.Tensor, key: torch.Tensor,
value: torch.Tensor, value: torch.Tensor,
attn_metadata: "XFormersMetadata", attn_metadata: "XFormersMetadata",
) -> torch.Tensor: ) -> torch.Tensor:
"""Use ixformer flash_attn_varlen_func for head_dim > 128. """Use ixformer flash_attn_varlen_func for head_dim > 128.
Verified on real BI-V100: flash_attn_func handles head_dim=256 Verified on real BI-V100: flash_attn_func handles head_dim=256
correctly (diff < 0.004, no NaN). For seq >= 1024, faster than correctly (diff < 0.004, no NaN). For seq >= 1024, faster than
PyTorch matmul. For profiling, sequences can be 20K+ tokens — this PyTorch matmul. For profiling, sequences can be 20K+ tokens — this
is dramatically faster than the previous Python Q-tiling fallback. is dramatically faster than the previous Python Q-tiling fallback.
Falls back to pure-math if flash_attn is unavailable. Falls back to pure-math if flash_attn is unavailable.
""" """
import ixformer as _ixf import ixformer as _ixf
assert attn_metadata.seq_lens is not None assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype orig_dtype = query.dtype
num_seqs = len(attn_metadata.seq_lens) num_seqs = len(attn_metadata.seq_lens)
q_flat = query.squeeze(0) # [T, H, D] q_flat = query.squeeze(0) # [T, H, D]
k_flat = key.squeeze(0) # [T, Hkv, D] k_flat = key.squeeze(0) # [T, Hkv, D]
v_flat = value.squeeze(0) v_flat = value.squeeze(0)
# Build cu_seqlens from seq_lens # Build cu_seqlens from seq_lens
seq_lens_list = list(attn_metadata.seq_lens) seq_lens_list = list(attn_metadata.seq_lens)
cu_seqlens = torch.zeros(num_seqs + 1, dtype=torch.int32, cu_seqlens = torch.zeros(num_seqs + 1, dtype=torch.int32,
device=query.device) device=query.device)
for i, sl in enumerate(seq_lens_list): for i, sl in enumerate(seq_lens_list):
cu_seqlens[i + 1] = cu_seqlens[i] + sl cu_seqlens[i + 1] = cu_seqlens[i] + sl
max_seqlen = max(seq_lens_list) max_seqlen = max(seq_lens_list)
try: try:
# Skip flash_attn during profiling — OOMs on large dummy batch # Skip flash_attn during profiling — OOMs on large dummy batch
import os import os
if os.environ.get("BI100_IN_STARTUP_PROFILE") == "1": if os.environ.get("BI100_IN_STARTUP_PROFILE") == "1":
raise RuntimeError("skip flash_attn during profiling") raise RuntimeError("skip flash_attn during profiling")
out = _ixf.flash_attn_varlen_func( out = _ixf.flash_attn_varlen_func(
q_flat.to(torch.float16), q_flat.to(torch.float16),
k_flat.to(torch.float16), k_flat.to(torch.float16),
v_flat.to(torch.float16), v_flat.to(torch.float16),
cu_seqlens, cu_seqlens, cu_seqlens, cu_seqlens,
max_seqlen, max_seqlen, max_seqlen, max_seqlen,
causal=True, causal=True,
) )
return out.to(orig_dtype).unsqueeze(0) return out.to(orig_dtype).unsqueeze(0)
except Exception: except Exception:
pass pass
# Fallback: pure-math Q-tiling (original implementation) # Fallback: pure-math Q-tiling (original implementation)
_Q_CHUNK = 256 _Q_CHUNK = 256
# During profiling, skip expensive attention — return zeros. # During profiling, skip expensive attention — return zeros.
# Profiling only measures memory footprint, not output correctness. # Profiling only measures memory footprint, not output correctness.
if os.environ.get("BI100_IN_STARTUP_PROFILE") == "1": if os.environ.get("BI100_IN_STARTUP_PROFILE") == "1":
return torch.zeros_like(query) return torch.zeros_like(query)
if (attn_metadata.query_start_loc is not None if (attn_metadata.query_start_loc is not None
and len(attn_metadata.query_start_loc) == num_seqs + 1): and len(attn_metadata.query_start_loc) == num_seqs + 1):
q_lens = [ q_lens = [
int(attn_metadata.query_start_loc[i + 1].item()) - int(attn_metadata.query_start_loc[i + 1].item()) -
int(attn_metadata.query_start_loc[i].item()) int(attn_metadata.query_start_loc[i].item())
for i in range(num_seqs) for i in range(num_seqs)
] ]
else: else:
q_lens = seq_lens_list q_lens = seq_lens_list
output = torch.empty_like(q_flat) output = torch.empty_like(q_flat)
seq_start = 0 seq_start = 0
for q_len in q_lens: for q_len in q_lens:
seq_end = seq_start + q_len seq_end = seq_start + q_len
k_s = k_flat[seq_start:seq_end].permute(1, 0, 2).float() k_s = k_flat[seq_start:seq_end].permute(1, 0, 2).float()
v_s = v_flat[seq_start:seq_end].permute(1, 0, 2).float() v_s = v_flat[seq_start:seq_end].permute(1, 0, 2).float()
if k_s.shape[0] != self.num_heads: if k_s.shape[0] != self.num_heads:
n = self.num_heads // k_s.shape[0] n = self.num_heads // k_s.shape[0]
k_s = k_s.repeat_interleave(n, dim=0).contiguous() k_s = k_s.repeat_interleave(n, dim=0).contiguous()
v_s = v_s.repeat_interleave(n, dim=0).contiguous() v_s = v_s.repeat_interleave(n, dim=0).contiguous()
k_pos = torch.arange(q_len, device=query.device) k_pos = torch.arange(q_len, device=query.device)
for qc_start in range(0, q_len, _Q_CHUNK): for qc_start in range(0, q_len, _Q_CHUNK):
qc_end = min(qc_start + _Q_CHUNK, q_len) qc_end = min(qc_start + _Q_CHUNK, q_len)
q_c = q_flat[seq_start + qc_start:seq_start + qc_end] \ q_c = q_flat[seq_start + qc_start:seq_start + qc_end] \
.permute(1, 0, 2).float() .permute(1, 0, 2).float()
attn_w = torch.matmul(q_c, k_s.transpose(-2, -1)) * self.scale attn_w = torch.matmul(q_c, k_s.transpose(-2, -1)) * self.scale
qc_q_pos = torch.arange(qc_start, qc_end, device=query.device) qc_q_pos = torch.arange(qc_start, qc_end, device=query.device)
mask = k_pos.unsqueeze(0) > qc_q_pos.unsqueeze(1) mask = k_pos.unsqueeze(0) > qc_q_pos.unsqueeze(1)
attn_w = attn_w.masked_fill(mask.unsqueeze(0), float("-inf")) attn_w = attn_w.masked_fill(mask.unsqueeze(0), float("-inf"))
attn_w = torch.softmax(attn_w, dim=-1) attn_w = torch.softmax(attn_w, dim=-1)
out_c = torch.matmul(attn_w, v_s).to(orig_dtype) out_c = torch.matmul(attn_w, v_s).to(orig_dtype)
output[seq_start + qc_start:seq_start + qc_end] = ( output[seq_start + qc_start:seq_start + qc_end] = (
out_c.permute(1, 0, 2)) out_c.permute(1, 0, 2))
seq_start = seq_end seq_start = seq_end
return output.unsqueeze(0) return output.unsqueeze(0)
''' '''
OLD_XFORMER_BLOCK = """\ OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp() self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None: if self.alibi_slopes is None:
# Add the batch dimension. # Add the batch dimension.
query = query.unsqueeze(0) query = query.unsqueeze(0)
key = key.unsqueeze(0) key = key.unsqueeze(0)
value = value.unsqueeze(0) value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward( out = xops.memory_efficient_attention_forward(
query, query,
key, key,
value, value,
attn_bias=attn_bias[0], attn_bias=attn_bias[0],
p=0.0, p=0.0,
scale=self.scale, scale=self.scale,
op = self.attn_op op = self.attn_op
) )
return out.view_as(original_query)\ return out.view_as(original_query)\
""" """
NEW_XFORMER_BLOCK = """\ NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp() self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None: if self.alibi_slopes is None:
# Add the batch dimension. # Add the batch dimension.
query = query.unsqueeze(0) query = query.unsqueeze(0)
key = key.unsqueeze(0) key = key.unsqueeze(0)
value = value.unsqueeze(0) value = value.unsqueeze(0)
if self.head_size > 128: if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata) out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else: else:
out = xops.memory_efficient_attention_forward( out = xops.memory_efficient_attention_forward(
query, query,
key, key,
value, value,
attn_bias=attn_bias[0], attn_bias=attn_bias[0],
p=0.0, p=0.0,
scale=self.scale, scale=self.scale,
op=self.attn_op, op=self.attn_op,
) )
return out.view_as(original_query)\ return out.view_as(original_query)\
""" """
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward(" INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
_PREFIX_CALL_OLD_BLOCK = """\ _PREFIX_CALL_OLD_BLOCK = """\
@@ -367,8 +367,8 @@ def patch_file(path):
required=True, required=True,
already_contains=( already_contains=(
"is_causal_decoder=(attn_type == AttentionType.DECODER)")) "is_causal_decoder=(attn_type == AttentionType.DECODER)"))
def patch_arg_utils(path): def patch_arg_utils(path):
replace_once( replace_once(
path, path,
@@ -382,8 +382,8 @@ def patch_arg_utils(path):
_MM_PREFIX_NEW_BLOCK, _MM_PREFIX_NEW_BLOCK,
required=True, required=True,
already_contains="Keeping prefix caching enabled for the Qwen3.6") already_contains="Keeping prefix caching enabled for the Qwen3.6")
def patch_logits_processor(path): def patch_logits_processor(path):
replace_once( replace_once(
path, path,
@@ -402,27 +402,27 @@ def patch_outlines_json_grammar(path):
], ],
required=True, required=True,
already_contains="JSON_WS:") already_contains="JSON_WS:")
def main(): def main():
print("=== patch_xformers_sdpa_seq (sequential, pure-math) ===") print("=== patch_xformers_sdpa_seq (sequential, pure-math) ===")
print(f"Target: {XFORMERS_PATH}") print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH) patch_file(XFORMERS_PATH)
print("\n=== patch_arg_utils (disable chunked-prefill auto-enable) ===") print("\n=== patch_arg_utils (disable chunked-prefill auto-enable) ===")
print(f"Target: {ARG_UTILS_PATH}") print(f"Target: {ARG_UTILS_PATH}")
patch_arg_utils(ARG_UTILS_PATH) patch_arg_utils(ARG_UTILS_PATH)
print("\n=== patch_logits_processor (seq_groups=None guard for chunked prefill) ===") print("\n=== patch_logits_processor (seq_groups=None guard for chunked prefill) ===")
print(f"Target: {LOGITS_PROC_PATH}") print(f"Target: {LOGITS_PROC_PATH}")
patch_logits_processor(LOGITS_PROC_PATH) patch_logits_processor(LOGITS_PROC_PATH)
print("\n=== patch_outlines_json_grammar (reject raw control chars) ===") print("\n=== patch_outlines_json_grammar (reject raw control chars) ===")
print(f"Target: {OUTLINES_DECODING_PATH}") print(f"Target: {OUTLINES_DECODING_PATH}")
patch_outlines_json_grammar(OUTLINES_DECODING_PATH) patch_outlines_json_grammar(OUTLINES_DECODING_PATH)
print("\nDone.") print("\nDone.")
if __name__ == "__main__": if __name__ == "__main__":
main() main()

View File

@@ -1,22 +1,22 @@
""" """
策略:顺序(per-sequence)— F.scaled_dot_product_attention,可走硬件 kernel 策略:顺序(per-sequence)— F.scaled_dot_product_attention,可走硬件 kernel
============================================================================= =============================================================================
逐条序列调用 F.scaled_dot_product_attention,is_causal=False + 显式因果 mask。 逐条序列调用 F.scaled_dot_product_attention,is_causal=False + 显式因果 mask。
与 patch_xformers_sdpa_seq.py(纯 matmul)的区别: 与 patch_xformers_sdpa_seq.py(纯 matmul)的区别:
SDPA 可分发到 Flash Attention / mem-efficient attention kernel, SDPA 可分发到 Flash Attention / mem-efficient attention kernel,
而纯 matmul 固定走 cublas。 而纯 matmul 固定走 cublas。
硬件限制(BI-V100): 硬件限制(BI-V100):
cudnnFlashAttnForward 不支持 is_causal=True(直接报错)。 cudnnFlashAttnForward 不支持 is_causal=True(直接报错)。
必须使用 is_causal=False + 显式 additive causal mask。 必须使用 is_causal=False + 显式 additive causal mask。
每条序列单独构造上三角 -inf mask,peak 显存 = max(seq_len)² × dtype, 每条序列单独构造上三角 -inf mask,peak 显存 = max(seq_len)² × dtype,
比 batch 版的 total_tokens² 小得多。 比 batch 版的 total_tokens² 小得多。
与 batch_kernel 的对比: 与 batch_kernel 的对比:
seq_kernel: 显存小,peak = max_single_seq²;并发 prefill 串行排队 seq_kernel: 显存小,peak = max_single_seq²;并发 prefill 串行排队
batch_kernel: 显存大,peak = total_tokens²;并发 prefill 一次并行处理, batch_kernel: 显存大,peak = total_tokens²;并发 prefill 一次并行处理,
通过 --max-num-batched-tokens 控制 total_tokens 上限 通过 --max-num-batched-tokens 控制 total_tokens 上限
Deploy: Deploy:
python3 modified_scripts/patch_xformers_sdpa_seq_kernel.py python3 modified_scripts/patch_xformers_sdpa_seq_kernel.py
""" """
@@ -24,122 +24,122 @@ Deploy:
from patch_utils import package_root, replace_once from patch_utils import package_root, replace_once
XFORMERS_PATH = package_root("vllm") / "attention" / "backends" / "xformers.py" XFORMERS_PATH = package_root("vllm") / "attention" / "backends" / "xformers.py"
FALLBACK_METHOD = ''' FALLBACK_METHOD = '''
def _run_sdpa_fallback( def _run_sdpa_fallback(
self, self,
query: torch.Tensor, query: torch.Tensor,
key: torch.Tensor, key: torch.Tensor,
value: torch.Tensor, value: torch.Tensor,
attn_metadata: "XFormersMetadata", attn_metadata: "XFormersMetadata",
) -> torch.Tensor: ) -> torch.Tensor:
"""顺序 F.scaled_dot_product_attention fallback(可走硬件 kernel)。 """顺序 F.scaled_dot_product_attention fallback(可走硬件 kernel)。
逐条序列调用 SDPA,is_causal=False + 显式上三角 additive mask。 逐条序列调用 SDPA,is_causal=False + 显式上三角 additive mask。
cudnnFlashAttnForward 不支持 is_causal=True,必须用显式 mask。 cudnnFlashAttnForward 不支持 is_causal=True,必须用显式 mask。
逐序列构造 mask,peak 显存 = max(seq_len)² × dtype(远小于 batch 版)。 逐序列构造 mask,peak 显存 = max(seq_len)² × dtype(远小于 batch 版)。
Args: Args:
query : [1, total_prefill_tokens, num_heads, head_dim] query : [1, total_prefill_tokens, num_heads, head_dim]
key : [1, total_prefill_tokens, num_kv_heads, head_dim] key : [1, total_prefill_tokens, num_kv_heads, head_dim]
value : [1, total_prefill_tokens, num_kv_heads, head_dim] value : [1, total_prefill_tokens, num_kv_heads, head_dim]
Returns: Returns:
[1, total_prefill_tokens, num_heads, head_dim] [1, total_prefill_tokens, num_heads, head_dim]
""" """
import torch.nn.functional as F import torch.nn.functional as F
assert attn_metadata.seq_lens is not None assert attn_metadata.seq_lens is not None
orig_dtype = query.dtype orig_dtype = query.dtype
q_flat = query.squeeze(0) # [T, H, D] q_flat = query.squeeze(0) # [T, H, D]
k_flat = key.squeeze(0) # [T, Hkv, D] k_flat = key.squeeze(0) # [T, Hkv, D]
v_flat = value.squeeze(0) v_flat = value.squeeze(0)
output = torch.empty_like(q_flat) output = torch.empty_like(q_flat)
start = 0 start = 0
for seq_len in attn_metadata.seq_lens: for seq_len in attn_metadata.seq_lens:
end = start + seq_len end = start + seq_len
# [1, H, L, D] # [1, H, L, D]
q_s = q_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0) q_s = q_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0)
k_s = k_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0) k_s = k_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0)
v_s = v_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0) v_s = v_flat[start:end].permute(1, 0, 2).contiguous().unsqueeze(0)
# GQA:展开 KV heads # GQA:展开 KV heads
if k_s.shape[1] != q_s.shape[1]: if k_s.shape[1] != q_s.shape[1]:
n = q_s.shape[1] // k_s.shape[1] n = q_s.shape[1] // k_s.shape[1]
k_s = k_s.repeat_interleave(n, dim=1).contiguous() k_s = k_s.repeat_interleave(n, dim=1).contiguous()
v_s = v_s.repeat_interleave(n, dim=1).contiguous() v_s = v_s.repeat_interleave(n, dim=1).contiguous()
# 逐序列因果 mask [L, L],上三角 -inf # 逐序列因果 mask [L, L],上三角 -inf
causal_mask = torch.tril( causal_mask = torch.tril(
torch.zeros(seq_len, seq_len, dtype=orig_dtype, device=q_s.device) torch.zeros(seq_len, seq_len, dtype=orig_dtype, device=q_s.device)
) )
causal_mask = causal_mask.masked_fill( causal_mask = causal_mask.masked_fill(
torch.triu(torch.ones(seq_len, seq_len, dtype=torch.bool, torch.triu(torch.ones(seq_len, seq_len, dtype=torch.bool,
device=q_s.device), diagonal=1), device=q_s.device), diagonal=1),
float("-inf"), float("-inf"),
) )
# is_causal=False + 显式 mask,规避 cudnnFlashAttnForward 不支持 is_causal=True # is_causal=False + 显式 mask,规避 cudnnFlashAttnForward 不支持 is_causal=True
out_s = F.scaled_dot_product_attention( out_s = F.scaled_dot_product_attention(
q_s, k_s, v_s, q_s, k_s, v_s,
attn_mask=causal_mask, attn_mask=causal_mask,
dropout_p=0.0, dropout_p=0.0,
is_causal=False, is_causal=False,
scale=self.scale, scale=self.scale,
) )
# [1, H, L, D] → [L, H, D] # [1, H, L, D] → [L, H, D]
output[start:end] = out_s.squeeze(0).permute(1, 0, 2).to(orig_dtype) output[start:end] = out_s.squeeze(0).permute(1, 0, 2).to(orig_dtype)
start = end start = end
return output.unsqueeze(0) # [1, T, H, D] return output.unsqueeze(0) # [1, T, H, D]
''' '''
OLD_XFORMER_BLOCK = """\ OLD_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp() self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None: if self.alibi_slopes is None:
# Add the batch dimension. # Add the batch dimension.
query = query.unsqueeze(0) query = query.unsqueeze(0)
key = key.unsqueeze(0) key = key.unsqueeze(0)
value = value.unsqueeze(0) value = value.unsqueeze(0)
out = xops.memory_efficient_attention_forward( out = xops.memory_efficient_attention_forward(
query, query,
key, key,
value, value,
attn_bias=attn_bias[0], attn_bias=attn_bias[0],
p=0.0, p=0.0,
scale=self.scale, scale=self.scale,
op = self.attn_op op = self.attn_op
) )
return out.view_as(original_query)\ return out.view_as(original_query)\
""" """
NEW_XFORMER_BLOCK = """\ NEW_XFORMER_BLOCK = """\
self.attn_op = xops.fmha.flash.FwOp() self.attn_op = xops.fmha.flash.FwOp()
if self.alibi_slopes is None: if self.alibi_slopes is None:
# Add the batch dimension. # Add the batch dimension.
query = query.unsqueeze(0) query = query.unsqueeze(0)
key = key.unsqueeze(0) key = key.unsqueeze(0)
value = value.unsqueeze(0) value = value.unsqueeze(0)
if self.head_size > 128: if self.head_size > 128:
out = self._run_sdpa_fallback(query, key, value, attn_metadata) out = self._run_sdpa_fallback(query, key, value, attn_metadata)
else: else:
out = xops.memory_efficient_attention_forward( out = xops.memory_efficient_attention_forward(
query, query,
key, key,
value, value,
attn_bias=attn_bias[0], attn_bias=attn_bias[0],
p=0.0, p=0.0,
scale=self.scale, scale=self.scale,
op=self.attn_op, op=self.attn_op,
) )
return out.view_as(original_query)\ return out.view_as(original_query)\
""" """
INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward(" INJECT_ANCHOR = " def _run_memory_efficient_xformers_forward("
def patch_file(path): def patch_file(path):
replace_once( replace_once(
path, path,
@@ -153,14 +153,14 @@ def patch_file(path):
NEW_XFORMER_BLOCK, NEW_XFORMER_BLOCK,
required=True, required=True,
already_contains="out = self._run_sdpa_fallback(query, key, value, attn_metadata)") already_contains="out = self._run_sdpa_fallback(query, key, value, attn_metadata)")
def main(): def main():
print("=== patch_xformers_sdpa_seq_kernel (seq, F.sdpa + kernel dispatch) ===") print("=== patch_xformers_sdpa_seq_kernel (seq, F.sdpa + kernel dispatch) ===")
print(f"Target: {XFORMERS_PATH}") print(f"Target: {XFORMERS_PATH}")
patch_file(XFORMERS_PATH) patch_file(XFORMERS_PATH)
print("\nDone.") print("\nDone.")
if __name__ == "__main__": if __name__ == "__main__":
main() main()

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -1,3 +1,3 @@
from .configuration_qwen3_5 import Qwen3_5Config, Qwen3_5TextConfig, Qwen3_5VisionConfig from .configuration_qwen3_5 import Qwen3_5Config, Qwen3_5TextConfig, Qwen3_5VisionConfig
__all__ = ["Qwen3_5Config", "Qwen3_5TextConfig", "Qwen3_5VisionConfig"] __all__ = ["Qwen3_5Config", "Qwen3_5TextConfig", "Qwen3_5VisionConfig"]

View File

@@ -1,20 +1,20 @@
# Adapted from transformers 5.2.0 for compatibility with transformers 4.55.3 + torch 2.1.0 # Adapted from transformers 5.2.0 for compatibility with transformers 4.55.3 + torch 2.1.0
# Stubs layer_type_validation and RopeParameters which do not exist in 4.55.3 # Stubs layer_type_validation and RopeParameters which do not exist in 4.55.3
import os import os
from typing import Optional, List from typing import Optional, List
from ...configuration_utils import PretrainedConfig as PreTrainedConfig from ...configuration_utils import PretrainedConfig as PreTrainedConfig
# --- Local stubs for APIs not present in transformers 4.55.3 --- # --- Local stubs for APIs not present in transformers 4.55.3 ---
# Always use these definitions; do NOT import from the older transformers # Always use these definitions; do NOT import from the older transformers
# as same-named functions there have incompatible signatures. # as same-named functions there have incompatible signatures.
def layer_type_validation(layer_types, num_hidden_layers=None, attention=True): def layer_type_validation(layer_types, num_hidden_layers=None, attention=True):
allowed = {"full_attention", "linear_attention"} allowed = {"full_attention", "linear_attention"}
if not all(lt in allowed for lt in layer_types): if not all(lt in allowed for lt in layer_types):
raise ValueError(f"layer_types entries must be in {allowed}, got {layer_types}") raise ValueError(f"layer_types entries must be in {allowed}, got {layer_types}")
if num_hidden_layers is not None and num_hidden_layers != len(layer_types): if num_hidden_layers is not None and num_hidden_layers != len(layer_types):
raise ValueError( raise ValueError(
f"num_hidden_layers ({num_hidden_layers}) != len(layer_types) ({len(layer_types)})" f"num_hidden_layers ({num_hidden_layers}) != len(layer_types) ({len(layer_types)})"
) )
@@ -57,7 +57,7 @@ def _vllm_layers_block_type(
"attention" if layer_type == "full_attention" else layer_type "attention" if layer_type == "full_attention" else layer_type
for layer_type in layer_types for layer_type in layer_types
] ]
try: try:
from typing import TypedDict from typing import TypedDict
except ImportError: except ImportError:
@@ -68,161 +68,161 @@ else:
rope_type: str rope_type: str
partial_rotary_factor: float partial_rotary_factor: float
factor: float factor: float
# --- End stubs --- # --- End stubs ---
class Qwen3_5TextConfig(PreTrainedConfig): class Qwen3_5TextConfig(PreTrainedConfig):
r""" r"""
Configuration for the text backbone of Qwen3.5 / Qwen3.6-35B-A3B models. Configuration for the text backbone of Qwen3.5 / Qwen3.6-35B-A3B models.
model_type is "qwen3_5_text" (used internally by the nested config). model_type is "qwen3_5_text" (used internally by the nested config).
""" """
model_type = "qwen3_5_text" model_type = "qwen3_5_text"
keys_to_ignore_at_inference = ["past_key_values"] keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
self, self,
vocab_size=248320, vocab_size=248320,
hidden_size=4096, hidden_size=4096,
intermediate_size=12288, intermediate_size=12288,
num_hidden_layers=32, num_hidden_layers=32,
num_attention_heads=16, num_attention_heads=16,
num_key_value_heads=4, num_key_value_heads=4,
hidden_act="silu", hidden_act="silu",
max_position_embeddings=32768, max_position_embeddings=32768,
initializer_range=0.02, initializer_range=0.02,
rms_norm_eps=1e-6, rms_norm_eps=1e-6,
use_cache=True, use_cache=True,
tie_word_embeddings=False, tie_word_embeddings=False,
rope_parameters=None, rope_parameters=None,
attention_bias=False, attention_bias=False,
attention_dropout=0.0, attention_dropout=0.0,
head_dim=256, head_dim=256,
linear_conv_kernel_dim=4, linear_conv_kernel_dim=4,
linear_key_head_dim=128, linear_key_head_dim=128,
linear_value_head_dim=128, linear_value_head_dim=128,
linear_num_key_heads=16, linear_num_key_heads=16,
linear_num_value_heads=32, linear_num_value_heads=32,
layer_types=None, layer_types=None,
pad_token_id=None, pad_token_id=None,
bos_token_id=None, bos_token_id=None,
eos_token_id=None, eos_token_id=None,
**kwargs, **kwargs,
): ):
self.pad_token_id = pad_token_id self.pad_token_id = pad_token_id
self.bos_token_id = bos_token_id self.bos_token_id = bos_token_id
self.eos_token_id = eos_token_id self.eos_token_id = eos_token_id
self.tie_word_embeddings = tie_word_embeddings self.tie_word_embeddings = tie_word_embeddings
self.vocab_size = vocab_size self.vocab_size = vocab_size
self.max_position_embeddings = max_position_embeddings self.max_position_embeddings = max_position_embeddings
self.hidden_size = hidden_size self.hidden_size = hidden_size
self.intermediate_size = intermediate_size self.intermediate_size = intermediate_size
self.num_hidden_layers = num_hidden_layers self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads self.num_attention_heads = num_attention_heads
self.num_key_value_heads = num_key_value_heads self.num_key_value_heads = num_key_value_heads
self.hidden_act = hidden_act self.hidden_act = hidden_act
self.initializer_range = initializer_range self.initializer_range = initializer_range
self.rms_norm_eps = rms_norm_eps self.rms_norm_eps = rms_norm_eps
self.use_cache = use_cache self.use_cache = use_cache
self.attention_bias = attention_bias self.attention_bias = attention_bias
self.attention_dropout = attention_dropout self.attention_dropout = attention_dropout
self.head_dim = head_dim self.head_dim = head_dim
self.rope_parameters = rope_parameters self.rope_parameters = rope_parameters
kwargs.setdefault("partial_rotary_factor", 0.25) kwargs.setdefault("partial_rotary_factor", 0.25)
self.layer_types = layer_types self.layer_types = layer_types
if self.layer_types is None: if self.layer_types is None:
interval_pattern = kwargs.get("full_attention_interval", 4) interval_pattern = kwargs.get("full_attention_interval", 4)
self.layer_types = [ self.layer_types = [
"linear_attention" if bool((i + 1) % interval_pattern) else "full_attention" "linear_attention" if bool((i + 1) % interval_pattern) else "full_attention"
for i in range(self.num_hidden_layers) for i in range(self.num_hidden_layers)
] ]
layer_type_validation(self.layer_types, self.num_hidden_layers) layer_type_validation(self.layer_types, self.num_hidden_layers)
self.linear_conv_kernel_dim = linear_conv_kernel_dim self.linear_conv_kernel_dim = linear_conv_kernel_dim
self.linear_key_head_dim = linear_key_head_dim self.linear_key_head_dim = linear_key_head_dim
self.linear_value_head_dim = linear_value_head_dim self.linear_value_head_dim = linear_value_head_dim
self.linear_num_key_heads = linear_num_key_heads self.linear_num_key_heads = linear_num_key_heads
self.linear_num_value_heads = linear_num_value_heads self.linear_num_value_heads = linear_num_value_heads
super().__init__(**kwargs) super().__init__(**kwargs)
class Qwen3_5VisionConfig(PreTrainedConfig): class Qwen3_5VisionConfig(PreTrainedConfig):
model_type = "qwen3_5_vision" model_type = "qwen3_5_vision"
def __init__( def __init__(
self, self,
depth=27, depth=27,
hidden_size=1152, hidden_size=1152,
hidden_act="gelu_pytorch_tanh", hidden_act="gelu_pytorch_tanh",
intermediate_size=4304, intermediate_size=4304,
num_heads=16, num_heads=16,
in_channels=3, in_channels=3,
patch_size=16, patch_size=16,
spatial_merge_size=2, spatial_merge_size=2,
temporal_patch_size=2, temporal_patch_size=2,
out_hidden_size=3584, out_hidden_size=3584,
num_position_embeddings=2304, num_position_embeddings=2304,
initializer_range=0.02, initializer_range=0.02,
**kwargs, **kwargs,
): ):
super().__init__(**kwargs) super().__init__(**kwargs)
self.depth = depth self.depth = depth
self.hidden_size = hidden_size self.hidden_size = hidden_size
self.hidden_act = hidden_act self.hidden_act = hidden_act
self.intermediate_size = intermediate_size self.intermediate_size = intermediate_size
self.num_heads = num_heads self.num_heads = num_heads
self.in_channels = in_channels self.in_channels = in_channels
self.patch_size = patch_size self.patch_size = patch_size
self.spatial_merge_size = spatial_merge_size self.spatial_merge_size = spatial_merge_size
self.temporal_patch_size = temporal_patch_size self.temporal_patch_size = temporal_patch_size
self.out_hidden_size = out_hidden_size self.out_hidden_size = out_hidden_size
self.num_position_embeddings = num_position_embeddings self.num_position_embeddings = num_position_embeddings
self.initializer_range = initializer_range self.initializer_range = initializer_range
class Qwen3_5Config(PreTrainedConfig): class Qwen3_5Config(PreTrainedConfig):
r""" r"""
Top-level configuration for Qwen3.5 / Qwen3.6-35B-A3B. Top-level configuration for Qwen3.5 / Qwen3.6-35B-A3B.
model_type = "qwen3_5" matches the model card / config.json. model_type = "qwen3_5" matches the model card / config.json.
Wraps Qwen3_5TextConfig (and optionally Qwen3_5VisionConfig for multimodal use). Wraps Qwen3_5TextConfig (and optionally Qwen3_5VisionConfig for multimodal use).
For vLLM text-only inference only text_config is consumed. For vLLM text-only inference only text_config is consumed.
""" """
model_type = "qwen3_5" model_type = "qwen3_5"
keys_to_ignore_at_inference = ["past_key_values"] keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
self, self,
text_config=None, text_config=None,
vision_config=None, vision_config=None,
image_token_id=248056, image_token_id=248056,
video_token_id=248057, video_token_id=248057,
vision_start_token_id=248053, vision_start_token_id=248053,
vision_end_token_id=248054, vision_end_token_id=248054,
tie_word_embeddings=False, tie_word_embeddings=False,
**kwargs, **kwargs,
): ):
serialized_mode = kwargs.pop(HYBRID_KV_ACCOUNTING_CONFIG, None) serialized_mode = kwargs.pop(HYBRID_KV_ACCOUNTING_CONFIG, None)
serialized_layers = kwargs.pop("layers_block_type", None) serialized_layers = kwargs.pop("layers_block_type", None)
if isinstance(text_config, dict): if isinstance(text_config, dict):
self.text_config = Qwen3_5TextConfig(**text_config) self.text_config = Qwen3_5TextConfig(**text_config)
elif text_config is None: elif text_config is None:
self.text_config = Qwen3_5TextConfig() self.text_config = Qwen3_5TextConfig()
else: else:
self.text_config = text_config self.text_config = text_config
if isinstance(vision_config, dict): if isinstance(vision_config, dict):
self.vision_config = Qwen3_5VisionConfig(**vision_config) self.vision_config = Qwen3_5VisionConfig(**vision_config)
elif vision_config is None: elif vision_config is None:
self.vision_config = Qwen3_5VisionConfig() self.vision_config = Qwen3_5VisionConfig()
else: else:
self.vision_config = vision_config self.vision_config = vision_config
self.image_token_id = image_token_id self.image_token_id = image_token_id
self.video_token_id = video_token_id self.video_token_id = video_token_id
self.vision_start_token_id = vision_start_token_id self.vision_start_token_id = vision_start_token_id
self.vision_end_token_id = vision_end_token_id self.vision_end_token_id = vision_end_token_id
self.tie_word_embeddings = tie_word_embeddings self.tie_word_embeddings = tie_word_embeddings
super().__init__(**kwargs) super().__init__(**kwargs)
@@ -237,6 +237,6 @@ class Qwen3_5Config(PreTrainedConfig):
f"{HYBRID_KV_ACCOUNTING_CONFIG}={mode!r}") f"{HYBRID_KV_ACCOUNTING_CONFIG}={mode!r}")
setattr(self, HYBRID_KV_ACCOUNTING_CONFIG, mode) setattr(self, HYBRID_KV_ACCOUNTING_CONFIG, mode)
self.layers_block_type = layers_block_type self.layers_block_type = layers_block_type
__all__ = ["Qwen3_5Config", "Qwen3_5TextConfig", "Qwen3_5VisionConfig"] __all__ = ["Qwen3_5Config", "Qwen3_5TextConfig", "Qwen3_5VisionConfig"]

View File

@@ -1,3 +1,3 @@
from .configuration_qwen3_5_moe import Qwen3_5MoeConfig, Qwen3_5MoeTextConfig from .configuration_qwen3_5_moe import Qwen3_5MoeConfig, Qwen3_5MoeTextConfig
__all__ = ["Qwen3_5MoeConfig", "Qwen3_5MoeTextConfig"] __all__ = ["Qwen3_5MoeConfig", "Qwen3_5MoeTextConfig"]

View File

@@ -1,20 +1,20 @@
# Adapted from transformers 5.2.0 for compatibility with transformers 4.55.3 + torch 2.1.0 # Adapted from transformers 5.2.0 for compatibility with transformers 4.55.3 + torch 2.1.0
# Source: transformers/models/qwen3_5_moe/configuration_qwen3_5_moe.py # Source: transformers/models/qwen3_5_moe/configuration_qwen3_5_moe.py
# Stubs layer_type_validation and RopeParameters which do not exist in 4.55.3 # Stubs layer_type_validation and RopeParameters which do not exist in 4.55.3
# Removes ignore_keys_at_rope_validation / base_model_tp_plan / base_model_pp_plan # Removes ignore_keys_at_rope_validation / base_model_tp_plan / base_model_pp_plan
# which are 5.x-only and irrelevant for vLLM inference. # which are 5.x-only and irrelevant for vLLM inference.
import os import os
from typing import Optional from typing import Optional
from ...configuration_utils import PretrainedConfig as PreTrainedConfig from ...configuration_utils import PretrainedConfig as PreTrainedConfig
# --- Local stubs for APIs not present in transformers 4.55.3 --- # --- Local stubs for APIs not present in transformers 4.55.3 ---
def layer_type_validation(layer_types, num_hidden_layers=None, attention=True): def layer_type_validation(layer_types, num_hidden_layers=None, attention=True):
allowed = {"full_attention", "linear_attention"} allowed = {"full_attention", "linear_attention"}
if not all(lt in allowed for lt in layer_types): if not all(lt in allowed for lt in layer_types):
raise ValueError(f"layer_types entries must be in {allowed}, got {layer_types}") raise ValueError(f"layer_types entries must be in {allowed}, got {layer_types}")
if num_hidden_layers is not None and num_hidden_layers != len(layer_types): if num_hidden_layers is not None and num_hidden_layers != len(layer_types):
raise ValueError( raise ValueError(
f"num_hidden_layers ({num_hidden_layers}) != len(layer_types) ({len(layer_types)})" f"num_hidden_layers ({num_hidden_layers}) != len(layer_types) ({len(layer_types)})"
) )
@@ -57,7 +57,7 @@ def _vllm_layers_block_type(
"attention" if layer_type == "full_attention" else layer_type "attention" if layer_type == "full_attention" else layer_type
for layer_type in layer_types for layer_type in layer_types
] ]
try: try:
from typing import TypedDict from typing import TypedDict
except ImportError: except ImportError:
@@ -68,171 +68,171 @@ else:
rope_type: str rope_type: str
partial_rotary_factor: float partial_rotary_factor: float
factor: float factor: float
# --- End stubs --- # --- End stubs ---
class Qwen3_5MoeTextConfig(PreTrainedConfig): class Qwen3_5MoeTextConfig(PreTrainedConfig):
r""" r"""
Configuration for the text backbone of Qwen3.5-MoE / Qwen3.6-35B-A3B models. Configuration for the text backbone of Qwen3.5-MoE / Qwen3.6-35B-A3B models.
model_type is "qwen3_5_moe_text" (used internally by the nested config). model_type is "qwen3_5_moe_text" (used internally by the nested config).
""" """
model_type = "qwen3_5_moe_text" model_type = "qwen3_5_moe_text"
keys_to_ignore_at_inference = ["past_key_values"] keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
self, self,
vocab_size=248320, vocab_size=248320,
hidden_size=2048, hidden_size=2048,
num_hidden_layers=40, num_hidden_layers=40,
num_attention_heads=16, num_attention_heads=16,
num_key_value_heads=2, num_key_value_heads=2,
hidden_act="silu", hidden_act="silu",
max_position_embeddings=32768, max_position_embeddings=32768,
initializer_range=0.02, initializer_range=0.02,
rms_norm_eps=1e-6, rms_norm_eps=1e-6,
use_cache=True, use_cache=True,
tie_word_embeddings=False, tie_word_embeddings=False,
rope_parameters=None, rope_parameters=None,
attention_bias=False, attention_bias=False,
attention_dropout=0.0, attention_dropout=0.0,
head_dim=256, head_dim=256,
linear_conv_kernel_dim=4, linear_conv_kernel_dim=4,
linear_key_head_dim=128, linear_key_head_dim=128,
linear_value_head_dim=128, linear_value_head_dim=128,
linear_num_key_heads=16, linear_num_key_heads=16,
linear_num_value_heads=32, linear_num_value_heads=32,
moe_intermediate_size=512, moe_intermediate_size=512,
shared_expert_intermediate_size=512, shared_expert_intermediate_size=512,
num_experts_per_tok=8, num_experts_per_tok=8,
num_experts=256, num_experts=256,
output_router_logits=False, output_router_logits=False,
router_aux_loss_coef=0.001, router_aux_loss_coef=0.001,
layer_types=None, layer_types=None,
pad_token_id=None, pad_token_id=None,
bos_token_id=None, bos_token_id=None,
eos_token_id=None, eos_token_id=None,
**kwargs, **kwargs,
): ):
self.pad_token_id = pad_token_id self.pad_token_id = pad_token_id
self.bos_token_id = bos_token_id self.bos_token_id = bos_token_id
self.eos_token_id = eos_token_id self.eos_token_id = eos_token_id
self.tie_word_embeddings = tie_word_embeddings self.tie_word_embeddings = tie_word_embeddings
self.vocab_size = vocab_size self.vocab_size = vocab_size
self.max_position_embeddings = max_position_embeddings self.max_position_embeddings = max_position_embeddings
self.hidden_size = hidden_size self.hidden_size = hidden_size
self.num_hidden_layers = num_hidden_layers self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads self.num_attention_heads = num_attention_heads
self.num_key_value_heads = num_key_value_heads self.num_key_value_heads = num_key_value_heads
self.hidden_act = hidden_act self.hidden_act = hidden_act
self.initializer_range = initializer_range self.initializer_range = initializer_range
self.rms_norm_eps = rms_norm_eps self.rms_norm_eps = rms_norm_eps
self.use_cache = use_cache self.use_cache = use_cache
self.attention_bias = attention_bias self.attention_bias = attention_bias
self.attention_dropout = attention_dropout self.attention_dropout = attention_dropout
self.head_dim = head_dim self.head_dim = head_dim
self.rope_parameters = rope_parameters self.rope_parameters = rope_parameters
kwargs.setdefault("partial_rotary_factor", 0.25) kwargs.setdefault("partial_rotary_factor", 0.25)
self.layer_types = layer_types self.layer_types = layer_types
if self.layer_types is None: if self.layer_types is None:
interval_pattern = kwargs.get("full_attention_interval", 4) interval_pattern = kwargs.get("full_attention_interval", 4)
self.layer_types = [ self.layer_types = [
"linear_attention" if bool((i + 1) % interval_pattern) else "full_attention" "linear_attention" if bool((i + 1) % interval_pattern) else "full_attention"
for i in range(self.num_hidden_layers) for i in range(self.num_hidden_layers)
] ]
layer_type_validation(self.layer_types, self.num_hidden_layers) layer_type_validation(self.layer_types, self.num_hidden_layers)
self.linear_conv_kernel_dim = linear_conv_kernel_dim self.linear_conv_kernel_dim = linear_conv_kernel_dim
self.linear_key_head_dim = linear_key_head_dim self.linear_key_head_dim = linear_key_head_dim
self.linear_value_head_dim = linear_value_head_dim self.linear_value_head_dim = linear_value_head_dim
self.linear_num_key_heads = linear_num_key_heads self.linear_num_key_heads = linear_num_key_heads
self.linear_num_value_heads = linear_num_value_heads self.linear_num_value_heads = linear_num_value_heads
self.moe_intermediate_size = moe_intermediate_size self.moe_intermediate_size = moe_intermediate_size
self.shared_expert_intermediate_size = shared_expert_intermediate_size self.shared_expert_intermediate_size = shared_expert_intermediate_size
self.num_experts_per_tok = num_experts_per_tok self.num_experts_per_tok = num_experts_per_tok
self.num_experts = num_experts self.num_experts = num_experts
self.output_router_logits = output_router_logits self.output_router_logits = output_router_logits
self.router_aux_loss_coef = router_aux_loss_coef self.router_aux_loss_coef = router_aux_loss_coef
super().__init__(**kwargs) super().__init__(**kwargs)
class Qwen3_5MoeVisionConfig(PreTrainedConfig): class Qwen3_5MoeVisionConfig(PreTrainedConfig):
model_type = "qwen3_5_moe" model_type = "qwen3_5_moe"
def __init__( def __init__(
self, self,
depth=27, depth=27,
hidden_size=1152, hidden_size=1152,
hidden_act="gelu_pytorch_tanh", hidden_act="gelu_pytorch_tanh",
intermediate_size=4304, intermediate_size=4304,
num_heads=16, num_heads=16,
in_channels=3, in_channels=3,
patch_size=16, patch_size=16,
spatial_merge_size=2, spatial_merge_size=2,
temporal_patch_size=2, temporal_patch_size=2,
out_hidden_size=3584, out_hidden_size=3584,
num_position_embeddings=2304, num_position_embeddings=2304,
initializer_range=0.02, initializer_range=0.02,
**kwargs, **kwargs,
): ):
super().__init__(**kwargs) super().__init__(**kwargs)
self.depth = depth self.depth = depth
self.hidden_size = hidden_size self.hidden_size = hidden_size
self.hidden_act = hidden_act self.hidden_act = hidden_act
self.intermediate_size = intermediate_size self.intermediate_size = intermediate_size
self.num_heads = num_heads self.num_heads = num_heads
self.in_channels = in_channels self.in_channels = in_channels
self.patch_size = patch_size self.patch_size = patch_size
self.spatial_merge_size = spatial_merge_size self.spatial_merge_size = spatial_merge_size
self.temporal_patch_size = temporal_patch_size self.temporal_patch_size = temporal_patch_size
self.out_hidden_size = out_hidden_size self.out_hidden_size = out_hidden_size
self.num_position_embeddings = num_position_embeddings self.num_position_embeddings = num_position_embeddings
self.initializer_range = initializer_range self.initializer_range = initializer_range
class Qwen3_5MoeConfig(PreTrainedConfig): class Qwen3_5MoeConfig(PreTrainedConfig):
r""" r"""
Top-level configuration for Qwen3.5-MoE / Qwen3.6-35B-A3B. Top-level configuration for Qwen3.5-MoE / Qwen3.6-35B-A3B.
model_type = "qwen3_5_moe" matches the model card / config.json. model_type = "qwen3_5_moe" matches the model card / config.json.
Wraps Qwen3_5MoeTextConfig (and optionally Qwen3_5MoeVisionConfig). Wraps Qwen3_5MoeTextConfig (and optionally Qwen3_5MoeVisionConfig).
For vLLM text-only inference only text_config is consumed. For vLLM text-only inference only text_config is consumed.
""" """
model_type = "qwen3_5_moe" model_type = "qwen3_5_moe"
keys_to_ignore_at_inference = ["past_key_values"] keys_to_ignore_at_inference = ["past_key_values"]
def __init__( def __init__(
self, self,
text_config=None, text_config=None,
vision_config=None, vision_config=None,
image_token_id=248056, image_token_id=248056,
video_token_id=248057, video_token_id=248057,
vision_start_token_id=248053, vision_start_token_id=248053,
vision_end_token_id=248054, vision_end_token_id=248054,
tie_word_embeddings=False, tie_word_embeddings=False,
**kwargs, **kwargs,
): ):
serialized_mode = kwargs.pop(HYBRID_KV_ACCOUNTING_CONFIG, None) serialized_mode = kwargs.pop(HYBRID_KV_ACCOUNTING_CONFIG, None)
serialized_layers = kwargs.pop("layers_block_type", None) serialized_layers = kwargs.pop("layers_block_type", None)
if isinstance(text_config, dict): if isinstance(text_config, dict):
self.text_config = Qwen3_5MoeTextConfig(**text_config) self.text_config = Qwen3_5MoeTextConfig(**text_config)
elif text_config is None: elif text_config is None:
self.text_config = Qwen3_5MoeTextConfig() self.text_config = Qwen3_5MoeTextConfig()
else: else:
self.text_config = text_config self.text_config = text_config
if isinstance(vision_config, dict): if isinstance(vision_config, dict):
self.vision_config = Qwen3_5MoeVisionConfig(**vision_config) self.vision_config = Qwen3_5MoeVisionConfig(**vision_config)
elif vision_config is None: elif vision_config is None:
self.vision_config = Qwen3_5MoeVisionConfig() self.vision_config = Qwen3_5MoeVisionConfig()
else: else:
self.vision_config = vision_config self.vision_config = vision_config
self.image_token_id = image_token_id self.image_token_id = image_token_id
self.video_token_id = video_token_id self.video_token_id = video_token_id
self.vision_start_token_id = vision_start_token_id self.vision_start_token_id = vision_start_token_id
self.vision_end_token_id = vision_end_token_id self.vision_end_token_id = vision_end_token_id
self.tie_word_embeddings = tie_word_embeddings self.tie_word_embeddings = tie_word_embeddings
super().__init__(**kwargs) super().__init__(**kwargs)
@@ -247,6 +247,6 @@ class Qwen3_5MoeConfig(PreTrainedConfig):
f"{HYBRID_KV_ACCOUNTING_CONFIG}={mode!r}") f"{HYBRID_KV_ACCOUNTING_CONFIG}={mode!r}")
setattr(self, HYBRID_KV_ACCOUNTING_CONFIG, mode) setattr(self, HYBRID_KV_ACCOUNTING_CONFIG, mode)
self.layers_block_type = layers_block_type self.layers_block_type = layers_block_type
__all__ = ["Qwen3_5MoeConfig", "Qwen3_5MoeTextConfig"] __all__ = ["Qwen3_5MoeConfig", "Qwen3_5MoeTextConfig"]

File diff suppressed because it is too large Load Diff

View File

@@ -1,16 +1,16 @@
""" """
Reasoning parser module for vLLM 0.6.3 (BI-V100 / Qwen3.6-35B-A3B adaptation). Reasoning parser module for vLLM 0.6.3 (BI-V100 / Qwen3.6-35B-A3B adaptation).
Usage: --reasoning-parser qwen3 Usage: --reasoning-parser qwen3
""" """
from vllm.reasoning.abs_reasoning_parsers import ReasoningParser, ReasoningParserManager from vllm.reasoning.abs_reasoning_parsers import ReasoningParser, ReasoningParserManager
__all__ = ["ReasoningParser", "ReasoningParserManager"] __all__ = ["ReasoningParser", "ReasoningParserManager"]
# Lazy-register Qwen3 parser; imported on first get_reasoning_parser("qwen3"). # Lazy-register Qwen3 parser; imported on first get_reasoning_parser("qwen3").
ReasoningParserManager.register_lazy( ReasoningParserManager.register_lazy(
"qwen3", "qwen3",
"vllm.reasoning.qwen3_reasoning_parser", "vllm.reasoning.qwen3_reasoning_parser",
"Qwen3ReasoningParser", "Qwen3ReasoningParser",
) )

View File

@@ -1,243 +1,243 @@
""" """
Abstract reasoning parser base classes for vLLM 0.6.3. Abstract reasoning parser base classes for vLLM 0.6.3.
Adapted from vllm-original/vllm/reasoning/abs_reasoning_parsers.py: Adapted from vllm-original/vllm/reasoning/abs_reasoning_parsers.py:
- Removed vllm.entrypoints.mcp, vllm.utils.collection_utils, import_utils - Removed vllm.entrypoints.mcp, vllm.utils.collection_utils, import_utils
- DeltaMessage from vllm 0.6.3 protocol path - DeltaMessage from vllm 0.6.3 protocol path
- TokenizerLike -> AnyTokenizer - TokenizerLike -> AnyTokenizer
- ReasoningParserManager: simplified eager + lazy registration - ReasoningParserManager: simplified eager + lazy registration
""" """
import importlib import importlib
from abc import abstractmethod from abc import abstractmethod
from collections.abc import Iterable, Sequence from collections.abc import Iterable, Sequence
from functools import cached_property from functools import cached_property
from typing import Any, Optional, TYPE_CHECKING from typing import Any, Optional, TYPE_CHECKING
if TYPE_CHECKING: if TYPE_CHECKING:
from vllm.entrypoints.openai.protocol import DeltaMessage from vllm.entrypoints.openai.protocol import DeltaMessage
from vllm.transformers_utils.tokenizer import AnyTokenizer from vllm.transformers_utils.tokenizer import AnyTokenizer
else: else:
DeltaMessage = Any DeltaMessage = Any
AnyTokenizer = Any AnyTokenizer = Any
class ReasoningParser: class ReasoningParser:
"""Abstract base for all reasoning parsers.""" """Abstract base for all reasoning parsers."""
def __init__(self, tokenizer: "AnyTokenizer", *args, **kwargs): def __init__(self, tokenizer: "AnyTokenizer", *args, **kwargs):
self.model_tokenizer = tokenizer self.model_tokenizer = tokenizer
@cached_property @cached_property
def vocab(self) -> dict: def vocab(self) -> dict:
return self.model_tokenizer.get_vocab() return self.model_tokenizer.get_vocab()
@abstractmethod @abstractmethod
def is_reasoning_end(self, input_ids: Sequence[int]) -> bool: def is_reasoning_end(self, input_ids: Sequence[int]) -> bool:
"""Return True once the reasoning block has closed in input_ids.""" """Return True once the reasoning block has closed in input_ids."""
def is_reasoning_end_streaming( def is_reasoning_end_streaming(
self, input_ids: Sequence[int], delta_ids: Iterable[int] self, input_ids: Sequence[int], delta_ids: Iterable[int]
) -> bool: ) -> bool:
return self.is_reasoning_end(input_ids) return self.is_reasoning_end(input_ids)
@abstractmethod @abstractmethod
def extract_content_ids(self, input_ids: list) -> list: def extract_content_ids(self, input_ids: list) -> list:
"""Return token ids that belong to the content (post-reasoning) part.""" """Return token ids that belong to the content (post-reasoning) part."""
def count_reasoning_tokens(self, token_ids: Sequence[int]) -> int: def count_reasoning_tokens(self, token_ids: Sequence[int]) -> int:
return 0 return 0
@abstractmethod @abstractmethod
def extract_reasoning( def extract_reasoning(
self, model_output: str, request: Any self, model_output: str, request: Any
) -> "tuple[Optional[str], Optional[str]]": ) -> "tuple[Optional[str], Optional[str]]":
""" """
Split a complete model output into (reasoning_text, content_text). Split a complete model output into (reasoning_text, content_text).
Either part may be None. Either part may be None.
""" """
@abstractmethod @abstractmethod
def extract_reasoning_streaming( def extract_reasoning_streaming(
self, self,
previous_text: str, previous_text: str,
current_text: str, current_text: str,
delta_text: str, delta_text: str,
previous_token_ids: Sequence[int], previous_token_ids: Sequence[int],
current_token_ids: Sequence[int], current_token_ids: Sequence[int],
delta_token_ids: Sequence[int], delta_token_ids: Sequence[int],
) -> Optional["DeltaMessage"]: ) -> Optional["DeltaMessage"]:
""" """
Extract reasoning from a streaming delta. Extract reasoning from a streaming delta.
Returns a DeltaMessage with reasoning_content and/or content set, Returns a DeltaMessage with reasoning_content and/or content set,
or None if this delta should be suppressed (control token). or None if this delta should be suppressed (control token).
""" """
class BaseThinkingReasoningParser(ReasoningParser): class BaseThinkingReasoningParser(ReasoningParser):
""" """
Base for parsers that use <start_token>...</end_token> delimiters. Base for parsers that use <start_token>...</end_token> delimiters.
Subclasses define start_token / end_token properties. Subclasses define start_token / end_token properties.
""" """
@property @property
@abstractmethod @abstractmethod
def start_token(self) -> str: def start_token(self) -> str:
raise NotImplementedError raise NotImplementedError
@property @property
@abstractmethod @abstractmethod
def end_token(self) -> str: def end_token(self) -> str:
raise NotImplementedError raise NotImplementedError
def __init__(self, tokenizer: "AnyTokenizer", *args, **kwargs): def __init__(self, tokenizer: "AnyTokenizer", *args, **kwargs):
super().__init__(tokenizer, *args, **kwargs) super().__init__(tokenizer, *args, **kwargs)
if not self.model_tokenizer: if not self.model_tokenizer:
raise ValueError("Tokenizer must be passed to ReasoningParser.") raise ValueError("Tokenizer must be passed to ReasoningParser.")
if not self.start_token or not self.end_token: if not self.start_token or not self.end_token:
raise ValueError("start_token and end_token must be defined.") raise ValueError("start_token and end_token must be defined.")
self.start_token_id: Optional[int] = self.vocab.get(self.start_token) self.start_token_id: Optional[int] = self.vocab.get(self.start_token)
self.end_token_id: Optional[int] = self.vocab.get(self.end_token) self.end_token_id: Optional[int] = self.vocab.get(self.end_token)
if self.start_token_id is None or self.end_token_id is None: if self.start_token_id is None or self.end_token_id is None:
raise RuntimeError( raise RuntimeError(
f"{self.__class__.__name__}: could not find think tokens " f"{self.__class__.__name__}: could not find think tokens "
f"'{self.start_token}'/'{self.end_token}' in tokenizer vocab." f"'{self.start_token}'/'{self.end_token}' in tokenizer vocab."
) )
def is_reasoning_end(self, input_ids: Sequence[int]) -> bool: def is_reasoning_end(self, input_ids: Sequence[int]) -> bool:
for token_id in reversed(input_ids): for token_id in reversed(input_ids):
if token_id == self.start_token_id: if token_id == self.start_token_id:
return False return False
if token_id == self.end_token_id: if token_id == self.end_token_id:
return True return True
return False return False
def is_reasoning_end_streaming( def is_reasoning_end_streaming(
self, input_ids: Sequence[int], delta_ids: Iterable[int] self, input_ids: Sequence[int], delta_ids: Iterable[int]
) -> bool: ) -> bool:
return self.end_token_id in delta_ids return self.end_token_id in delta_ids
def extract_content_ids(self, input_ids: list) -> list: def extract_content_ids(self, input_ids: list) -> list:
if self.end_token_id not in input_ids[:-1]: if self.end_token_id not in input_ids[:-1]:
return [] return []
return input_ids[input_ids.index(self.end_token_id) + 1:] return input_ids[input_ids.index(self.end_token_id) + 1:]
def count_reasoning_tokens(self, token_ids: Sequence[int]) -> int: def count_reasoning_tokens(self, token_ids: Sequence[int]) -> int:
count = 0 count = 0
depth = 0 depth = 0
for tid in token_ids: for tid in token_ids:
if tid == self.start_token_id: if tid == self.start_token_id:
depth += 1 depth += 1
elif tid == self.end_token_id: elif tid == self.end_token_id:
if depth > 0: if depth > 0:
depth -= 1 depth -= 1
elif depth > 0: elif depth > 0:
count += 1 count += 1
return count return count
def extract_reasoning( def extract_reasoning(
self, model_output: str, request: Any self, model_output: str, request: Any
) -> "tuple[Optional[str], Optional[str]]": ) -> "tuple[Optional[str], Optional[str]]":
# Strip <think> if the model generated it (old-style template). # Strip <think> if the model generated it (old-style template).
parts = model_output.partition(self.start_token) parts = model_output.partition(self.start_token)
model_output = parts[2] if parts[1] else parts[0] model_output = parts[2] if parts[1] else parts[0]
if self.end_token not in model_output: if self.end_token not in model_output:
return model_output, None return model_output, None
reasoning, _, content = model_output.partition(self.end_token) reasoning, _, content = model_output.partition(self.end_token)
return reasoning, content or None return reasoning, content or None
def extract_reasoning_streaming( def extract_reasoning_streaming(
self, self,
previous_text: str, previous_text: str,
current_text: str, current_text: str,
delta_text: str, delta_text: str,
previous_token_ids: Sequence[int], previous_token_ids: Sequence[int],
current_token_ids: Sequence[int], current_token_ids: Sequence[int],
delta_token_ids: Sequence[int], delta_token_ids: Sequence[int],
) -> Optional["DeltaMessage"]: ) -> Optional["DeltaMessage"]:
from vllm.entrypoints.openai.protocol import DeltaMessage as _DeltaMessage from vllm.entrypoints.openai.protocol import DeltaMessage as _DeltaMessage
# Suppress lone control tokens. # Suppress lone control tokens.
if len(delta_token_ids) == 1 and delta_token_ids[0] in ( if len(delta_token_ids) == 1 and delta_token_ids[0] in (
self.start_token_id, self.end_token_id self.start_token_id, self.end_token_id
): ):
return None return None
start_in_prev = self.start_token_id in previous_token_ids start_in_prev = self.start_token_id in previous_token_ids
start_in_delta = self.start_token_id in delta_token_ids start_in_delta = self.start_token_id in delta_token_ids
end_in_prev = self.end_token_id in previous_token_ids end_in_prev = self.end_token_id in previous_token_ids
end_in_delta = self.end_token_id in delta_token_ids end_in_delta = self.end_token_id in delta_token_ids
if start_in_prev: if start_in_prev:
if end_in_delta: if end_in_delta:
end_idx = delta_text.find(self.end_token) end_idx = delta_text.find(self.end_token)
reasoning = delta_text[:end_idx] if end_idx >= 0 else "" reasoning = delta_text[:end_idx] if end_idx >= 0 else ""
content = delta_text[end_idx + len(self.end_token):] if end_idx >= 0 else None content = delta_text[end_idx + len(self.end_token):] if end_idx >= 0 else None
return _DeltaMessage( return _DeltaMessage(
reasoning_content=reasoning or None, reasoning_content=reasoning or None,
content=content or None, content=content or None,
) )
elif end_in_prev: elif end_in_prev:
return _DeltaMessage(content=delta_text) return _DeltaMessage(content=delta_text)
else: else:
return _DeltaMessage(reasoning_content=delta_text) return _DeltaMessage(reasoning_content=delta_text)
elif start_in_delta: elif start_in_delta:
if end_in_delta: if end_in_delta:
start_idx = delta_text.find(self.start_token) start_idx = delta_text.find(self.start_token)
end_idx = delta_text.find(self.end_token) end_idx = delta_text.find(self.end_token)
reasoning = delta_text[start_idx + len(self.start_token):end_idx] reasoning = delta_text[start_idx + len(self.start_token):end_idx]
content = delta_text[end_idx + len(self.end_token):] content = delta_text[end_idx + len(self.end_token):]
return _DeltaMessage( return _DeltaMessage(
reasoning_content=reasoning or None, reasoning_content=reasoning or None,
content=content or None, content=content or None,
) )
else: else:
return _DeltaMessage(reasoning_content=delta_text) return _DeltaMessage(reasoning_content=delta_text)
else: else:
return _DeltaMessage(content=delta_text) return _DeltaMessage(content=delta_text)
class ReasoningParserManager: class ReasoningParserManager:
""" """
Registry for ReasoningParser implementations. Registry for ReasoningParser implementations.
Supports eager and lazy registration. Supports eager and lazy registration.
""" """
_parsers: dict = {} # name -> class (eager) _parsers: dict = {} # name -> class (eager)
_lazy: dict = {} # name -> (module_path, class_name) _lazy: dict = {} # name -> (module_path, class_name)
@classmethod @classmethod
def register_module(cls, name: str, parser_cls: type) -> None: def register_module(cls, name: str, parser_cls: type) -> None:
"""Eagerly register a ReasoningParser class.""" """Eagerly register a ReasoningParser class."""
if not issubclass(parser_cls, ReasoningParser): if not issubclass(parser_cls, ReasoningParser):
raise TypeError(f"{parser_cls} is not a ReasoningParser subclass.") raise TypeError(f"{parser_cls} is not a ReasoningParser subclass.")
cls._parsers[name] = parser_cls cls._parsers[name] = parser_cls
@classmethod @classmethod
def register_lazy(cls, name: str, module_path: str, class_name: str) -> None: def register_lazy(cls, name: str, module_path: str, class_name: str) -> None:
"""Register a parser for deferred import.""" """Register a parser for deferred import."""
cls._lazy[name] = (module_path, class_name) cls._lazy[name] = (module_path, class_name)
@classmethod @classmethod
def get_reasoning_parser(cls, name: str) -> type: def get_reasoning_parser(cls, name: str) -> type:
if name in cls._parsers: if name in cls._parsers:
return cls._parsers[name] return cls._parsers[name]
if name in cls._lazy: if name in cls._lazy:
module_path, class_name = cls._lazy[name] module_path, class_name = cls._lazy[name]
mod = importlib.import_module(module_path) mod = importlib.import_module(module_path)
parser_cls = getattr(mod, class_name) parser_cls = getattr(mod, class_name)
cls._parsers[name] = parser_cls cls._parsers[name] = parser_cls
return parser_cls return parser_cls
registered = sorted(set(cls._parsers) | set(cls._lazy)) registered = sorted(set(cls._parsers) | set(cls._lazy))
raise KeyError( raise KeyError(
f"Reasoning parser '{name}' not found. " f"Reasoning parser '{name}' not found. "
f"Available: {registered}" f"Available: {registered}"
) )
@classmethod @classmethod
def list_registered(cls) -> list: def list_registered(cls) -> list:
return sorted(set(cls._parsers) | set(cls._lazy)) return sorted(set(cls._parsers) | set(cls._lazy))

View File

@@ -1,40 +1,40 @@
""" """
Reasoning parser for Qwen3 / Qwen3.5 / Qwen3.6 model family. Reasoning parser for Qwen3 / Qwen3.5 / Qwen3.6 model family.
Adapted from vllm-original/vllm/reasoning/qwen3_reasoning_parser.py. Adapted from vllm-original/vllm/reasoning/qwen3_reasoning_parser.py.
The model uses <think>...</think> to wrap chain-of-thought output. The model uses <think>...</think> to wrap chain-of-thought output.
For Qwen3.5+ the chat template injects <think> into the prompt, so only For Qwen3.5+ the chat template injects <think> into the prompt, so only
</think> appears in the generated tokens; older templates generate <think> </think> appears in the generated tokens; older templates generate <think>
themselves. Both styles are handled. themselves. Both styles are handled.
""" """
from typing import Optional, Sequence, Any from typing import Optional, Sequence, Any
from vllm.reasoning.abs_reasoning_parsers import ( from vllm.reasoning.abs_reasoning_parsers import (
BaseThinkingReasoningParser, BaseThinkingReasoningParser,
ReasoningParserManager, ReasoningParserManager,
) )
class Qwen3ReasoningParser(BaseThinkingReasoningParser): class Qwen3ReasoningParser(BaseThinkingReasoningParser):
def __init__(self, tokenizer: Any, *args, **kwargs): def __init__(self, tokenizer: Any, *args, **kwargs):
super().__init__(tokenizer, *args, **kwargs) super().__init__(tokenizer, *args, **kwargs)
chat_kwargs = kwargs.get("chat_template_kwargs", {}) or {} chat_kwargs = kwargs.get("chat_template_kwargs", {}) or {}
self.thinking_enabled = chat_kwargs.get("enable_thinking", True) self.thinking_enabled = chat_kwargs.get("enable_thinking", True)
@property @property
def start_token(self) -> str: def start_token(self) -> str:
return "<think>" return "<think>"
@property @property
def end_token(self) -> str: def end_token(self) -> str:
return "</think>" return "</think>"
def extract_reasoning( def extract_reasoning(
self, model_output: str, request: Any self, model_output: str, request: Any
) -> "tuple[Optional[str], Optional[str]]": ) -> "tuple[Optional[str], Optional[str]]":
# Strip <think> if the model generated it (old template / edge case). # Strip <think> if the model generated it (old template / edge case).
parts = model_output.partition(self.start_token) parts = model_output.partition(self.start_token)
model_output = parts[2] if parts[1] else parts[0] model_output = parts[2] if parts[1] else parts[0]
@@ -47,66 +47,66 @@ class Qwen3ReasoningParser(BaseThinkingReasoningParser):
if self.end_token not in model_output: if self.end_token not in model_output:
# Thinking enabled but output truncated before </think>. # Thinking enabled but output truncated before </think>.
return model_output, None return model_output, None
reasoning, _, content = model_output.partition(self.end_token) reasoning, _, content = model_output.partition(self.end_token)
return reasoning, content or None return reasoning, content or None
def count_reasoning_tokens(self, token_ids: Sequence[int]) -> int: def count_reasoning_tokens(self, token_ids: Sequence[int]) -> int:
token_ids = list(token_ids) token_ids = list(token_ids)
if self.start_token_id in token_ids: if self.start_token_id in token_ids:
# Old-style template: model generates <think> itself. # Old-style template: model generates <think> itself.
# Use depth-counting from the base class. # Use depth-counting from the base class.
return super().count_reasoning_tokens(token_ids) return super().count_reasoning_tokens(token_ids)
elif self.end_token_id in token_ids: elif self.end_token_id in token_ids:
# New-style template (Qwen3.5+): <think> is injected into the # New-style template (Qwen3.5+): <think> is injected into the
# prompt, so output starts already inside the thinking block. # prompt, so output starts already inside the thinking block.
# Every token before </think> is a reasoning token. # Every token before </think> is a reasoning token.
return token_ids.index(self.end_token_id) return token_ids.index(self.end_token_id)
else: else:
# No </think> in output: either truncated (all reasoning) # No </think> in output: either truncated (all reasoning)
# or thinking disabled (none). # or thinking disabled (none).
return len(token_ids) if self.thinking_enabled else 0 return len(token_ids) if self.thinking_enabled else 0
def extract_reasoning_streaming( def extract_reasoning_streaming(
self, self,
previous_text: str, previous_text: str,
current_text: str, current_text: str,
delta_text: str, delta_text: str,
previous_token_ids: Sequence[int], previous_token_ids: Sequence[int],
current_token_ids: Sequence[int], current_token_ids: Sequence[int],
delta_token_ids: Sequence[int], delta_token_ids: Sequence[int],
): ):
from vllm.entrypoints.openai.protocol import DeltaMessage from vllm.entrypoints.openai.protocol import DeltaMessage
if not self.thinking_enabled: if not self.thinking_enabled:
return DeltaMessage(content=delta_text) if delta_text else None return DeltaMessage(content=delta_text) if delta_text else None
# Strip <think> from delta if the model generates it itself. # Strip <think> from delta if the model generates it itself.
if self.start_token_id in delta_token_ids: if self.start_token_id in delta_token_ids:
start_idx = delta_text.find(self.start_token) start_idx = delta_text.find(self.start_token)
if start_idx >= 0: if start_idx >= 0:
delta_text = delta_text[start_idx + len(self.start_token):] delta_text = delta_text[start_idx + len(self.start_token):]
if self.end_token_id in delta_token_ids: if self.end_token_id in delta_token_ids:
end_idx = delta_text.find(self.end_token) end_idx = delta_text.find(self.end_token)
if end_idx >= 0: if end_idx >= 0:
reasoning = delta_text[:end_idx] reasoning = delta_text[:end_idx]
content = delta_text[end_idx + len(self.end_token):] content = delta_text[end_idx + len(self.end_token):]
if not reasoning and not content: if not reasoning and not content:
return None return None
return DeltaMessage( return DeltaMessage(
reasoning_content=reasoning or None, reasoning_content=reasoning or None,
content=content or None, content=content or None,
) )
return None return None
if not delta_text: if not delta_text:
return None return None
elif self.end_token_id in previous_token_ids: elif self.end_token_id in previous_token_ids:
return DeltaMessage(content=delta_text) return DeltaMessage(content=delta_text)
else: else:
return DeltaMessage(reasoning_content=delta_text) return DeltaMessage(reasoning_content=delta_text)
# Register immediately when this module is imported. # Register immediately when this module is imported.
ReasoningParserManager.register_module("qwen3", Qwen3ReasoningParser) ReasoningParserManager.register_module("qwen3", Qwen3ReasoningParser)

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -1,43 +1,43 @@
import math import math
from typing import List, Optional from typing import List, Optional
from vllm.core.block.common import BlockList from vllm.core.block.common import BlockList
from vllm.core.block.interfaces import Block, DeviceAwareBlockAllocator from vllm.core.block.interfaces import Block, DeviceAwareBlockAllocator
from vllm.utils import Device, cdiv, chunk_list from vllm.utils import Device, cdiv, chunk_list
class BlockTable: class BlockTable:
"""A class to manage blocks for a specific sequence. """A class to manage blocks for a specific sequence.
The BlockTable maps a sequence of tokens to a list of blocks, where each The BlockTable maps a sequence of tokens to a list of blocks, where each
block represents a contiguous memory allocation for a portion of the block represents a contiguous memory allocation for a portion of the
sequence. The blocks are managed by a DeviceAwareBlockAllocator, which is sequence. The blocks are managed by a DeviceAwareBlockAllocator, which is
responsible for allocating and freeing memory for the blocks. responsible for allocating and freeing memory for the blocks.
Args: Args:
block_size (int): The maximum number of tokens that can be stored in a block_size (int): The maximum number of tokens that can be stored in a
single block. single block.
block_allocator (DeviceAwareBlockAllocator): The block allocator used to block_allocator (DeviceAwareBlockAllocator): The block allocator used to
manage memory for the blocks. manage memory for the blocks.
_blocks (Optional[List[Block]], optional): An optional list of existing _blocks (Optional[List[Block]], optional): An optional list of existing
blocks to initialize the BlockTable with. If not provided, an empty blocks to initialize the BlockTable with. If not provided, an empty
BlockTable is created. BlockTable is created.
max_block_sliding_window (Optional[int], optional): The number of max_block_sliding_window (Optional[int], optional): The number of
blocks to keep around for each sequance. If None, all blocks blocks to keep around for each sequance. If None, all blocks
are kept (eg., when sliding window is not used). are kept (eg., when sliding window is not used).
It should at least fit the sliding window size of the model. It should at least fit the sliding window size of the model.
Attributes: Attributes:
_block_size (int): The maximum number of tokens that can be stored in a _block_size (int): The maximum number of tokens that can be stored in a
single block. single block.
_allocator (DeviceAwareBlockAllocator): The block allocator used to _allocator (DeviceAwareBlockAllocator): The block allocator used to
manage memory for the blocks. manage memory for the blocks.
_blocks (Optional[List[Block]]): The list of blocks managed by this _blocks (Optional[List[Block]]): The list of blocks managed by this
BlockTable. BlockTable.
_num_full_slots (int): The number of tokens currently stored in the _num_full_slots (int): The number of tokens currently stored in the
blocks. blocks.
""" """
def __init__( def __init__(
self, self,
block_size: int, block_size: int,
@@ -52,55 +52,55 @@ class BlockTable:
if _blocks is None: if _blocks is None:
_blocks = [] _blocks = []
self._blocks: BlockList = BlockList(_blocks) self._blocks: BlockList = BlockList(_blocks)
self._max_block_sliding_window = max_block_sliding_window self._max_block_sliding_window = max_block_sliding_window
self._num_full_slots = self._get_num_token_ids() self._num_full_slots = self._get_num_token_ids()
@staticmethod @staticmethod
def get_num_required_blocks(token_ids: List[int], def get_num_required_blocks(token_ids: List[int],
block_size: int, block_size: int,
num_lookahead_slots: int = 0) -> int: num_lookahead_slots: int = 0) -> int:
"""Calculates the minimum number of blocks required to store a given """Calculates the minimum number of blocks required to store a given
sequence of token IDs along with any look-ahead slots that may be sequence of token IDs along with any look-ahead slots that may be
required (like in multi-step + chunked-prefill). required (like in multi-step + chunked-prefill).
This assumes worst-case scenario, where every block requires a new This assumes worst-case scenario, where every block requires a new
allocation (e.g. ignoring prefix caching). allocation (e.g. ignoring prefix caching).
Args: Args:
token_ids (List[int]): The sequence of token IDs to be stored. token_ids (List[int]): The sequence of token IDs to be stored.
block_size (int): The maximum number of tokens that can be stored in block_size (int): The maximum number of tokens that can be stored in
a single block. a single block.
num_lookahead_slots (int): look-ahead slots that the sequence may num_lookahead_slots (int): look-ahead slots that the sequence may
require. require.
Returns: Returns:
int: The minimum number of blocks required to store the given int: The minimum number of blocks required to store the given
sequence of token IDs along with any required look-ahead slots. sequence of token IDs along with any required look-ahead slots.
""" """
return cdiv(len(token_ids) + num_lookahead_slots, block_size) return cdiv(len(token_ids) + num_lookahead_slots, block_size)
def allocate(self, def allocate(self,
token_ids: List[int], token_ids: List[int],
device: Device = Device.GPU) -> None: device: Device = Device.GPU) -> None:
"""Allocates memory blocks for storing the given sequence of token IDs. """Allocates memory blocks for storing the given sequence of token IDs.
This method allocates the required number of blocks to store the given This method allocates the required number of blocks to store the given
sequence of token IDs. sequence of token IDs.
Args: Args:
token_ids (List[int]): The sequence of token IDs to be stored. token_ids (List[int]): The sequence of token IDs to be stored.
device (Device, optional): The device on which the blocks should be device (Device, optional): The device on which the blocks should be
allocated. Defaults to Device.GPU. allocated. Defaults to Device.GPU.
""" """
assert not self._is_allocated assert not self._is_allocated
assert token_ids assert token_ids
blocks = self._allocate_blocks_for_token_ids(prev_block=None, blocks = self._allocate_blocks_for_token_ids(prev_block=None,
token_ids=token_ids, token_ids=token_ids,
device=device) device=device)
self.update(blocks) self.update(blocks)
self._num_full_slots = len(token_ids) self._num_full_slots = len(token_ids)
def update(self, blocks: List[Block]) -> None: def update(self, blocks: List[Block]) -> None:
"""Resets the table to the newly provided blocks """Resets the table to the newly provided blocks
(with their corresponding block ids) (with their corresponding block ids)
@@ -115,106 +115,106 @@ class BlockTable:
if block_hash is not None: if block_hash is not None:
content_hashes.append(block_hash) content_hashes.append(block_hash)
return content_hashes return content_hashes
def append_token_ids(self, def append_token_ids(self,
token_ids: List[int], token_ids: List[int],
num_lookahead_slots: int = 0, num_lookahead_slots: int = 0,
num_computed_slots: Optional[int] = None) -> None: num_computed_slots: Optional[int] = None) -> None:
"""Appends a sequence of token IDs to the existing blocks in the """Appends a sequence of token IDs to the existing blocks in the
BlockTable. BlockTable.
This method appends the given sequence of token IDs to the existing This method appends the given sequence of token IDs to the existing
blocks in the BlockTable. If there is not enough space in the existing blocks in the BlockTable. If there is not enough space in the existing
blocks, new blocks are allocated using the `ensure_num_empty_slots` blocks, new blocks are allocated using the `ensure_num_empty_slots`
method to accommodate the additional tokens. method to accommodate the additional tokens.
The token IDs are divided into chunks of size `block_size` (except for The token IDs are divided into chunks of size `block_size` (except for
the first chunk, which may be smaller), and each chunk is appended to a the first chunk, which may be smaller), and each chunk is appended to a
separate block. separate block.
Args: Args:
token_ids (List[int]): The sequence of token IDs to be appended. token_ids (List[int]): The sequence of token IDs to be appended.
num_computed_slots (Optional[int]): The number of KV cache slots num_computed_slots (Optional[int]): The number of KV cache slots
that are already filled (computed). that are already filled (computed).
When sliding window is enabled, this is used to compute how many When sliding window is enabled, this is used to compute how many
blocks to drop at the front of the sequence. blocks to drop at the front of the sequence.
Without sliding window, None can be passed. Without sliding window, None can be passed.
Without chunked prefill, it should be the same as Without chunked prefill, it should be the same as
_num_full_slots. _num_full_slots.
""" """
assert self._is_allocated, "no blocks have been allocated" assert self._is_allocated, "no blocks have been allocated"
assert len(self._blocks) > 0 assert len(self._blocks) > 0
# Drop blocks that are no longer needed due to sliding window # Drop blocks that are no longer needed due to sliding window
if self._max_block_sliding_window is not None: if self._max_block_sliding_window is not None:
null_block = self._allocator.allocate_or_get_null_block() null_block = self._allocator.allocate_or_get_null_block()
assert num_computed_slots is not None assert num_computed_slots is not None
end_block_idx = (num_computed_slots // end_block_idx = (num_computed_slots //
self._block_size) - self._max_block_sliding_window self._block_size) - self._max_block_sliding_window
for idx in range(0, end_block_idx): for idx in range(0, end_block_idx):
b = self._blocks[idx] b = self._blocks[idx]
if b is not null_block: if b is not null_block:
self._allocator.free(b) self._allocator.free(b)
self._blocks[idx] = null_block self._blocks[idx] = null_block
# Ensure there are enough empty slots for the new tokens plus # Ensure there are enough empty slots for the new tokens plus
# lookahead slots # lookahead slots
self.ensure_num_empty_slots(num_empty_slots=len(token_ids) + self.ensure_num_empty_slots(num_empty_slots=len(token_ids) +
num_lookahead_slots) num_lookahead_slots)
# Update the blocks with the new tokens # Update the blocks with the new tokens
first_block_idx = self._num_full_slots // self._block_size first_block_idx = self._num_full_slots // self._block_size
token_blocks = self._chunk_token_blocks_for_append(token_ids) token_blocks = self._chunk_token_blocks_for_append(token_ids)
for i, token_block in enumerate(token_blocks): for i, token_block in enumerate(token_blocks):
self._blocks.append_token_ids(first_block_idx + i, token_block) self._blocks.append_token_ids(first_block_idx + i, token_block)
self._num_full_slots += len(token_ids) self._num_full_slots += len(token_ids)
def ensure_num_empty_slots(self, num_empty_slots: int) -> None: def ensure_num_empty_slots(self, num_empty_slots: int) -> None:
"""Ensures that the BlockTable has at least the specified number of """Ensures that the BlockTable has at least the specified number of
empty slots available. empty slots available.
This method checks if the BlockTable has enough empty slots (i.e., This method checks if the BlockTable has enough empty slots (i.e.,
available space) to accommodate the requested number of tokens. If not, available space) to accommodate the requested number of tokens. If not,
it allocates additional blocks on the GPU to ensure that the required it allocates additional blocks on the GPU to ensure that the required
number of empty slots is available. number of empty slots is available.
Args: Args:
num_empty_slots (int): The minimum number of empty slots required. num_empty_slots (int): The minimum number of empty slots required.
""" """
# Currently the block table only supports # Currently the block table only supports
# appending tokens to GPU blocks. # appending tokens to GPU blocks.
device = Device.GPU device = Device.GPU
assert self._is_allocated assert self._is_allocated
if self._num_empty_slots >= num_empty_slots: if self._num_empty_slots >= num_empty_slots:
return return
slots_to_allocate = num_empty_slots - self._num_empty_slots slots_to_allocate = num_empty_slots - self._num_empty_slots
blocks_to_allocate = cdiv(slots_to_allocate, self._block_size) blocks_to_allocate = cdiv(slots_to_allocate, self._block_size)
for _ in range(blocks_to_allocate): for _ in range(blocks_to_allocate):
assert len(self._blocks) > 0 assert len(self._blocks) > 0
self._blocks.append( self._blocks.append(
self._allocator.allocate_mutable_block( self._allocator.allocate_mutable_block(
prev_block=self._blocks[-1], device=device)) prev_block=self._blocks[-1], device=device))
def fork(self) -> "BlockTable": def fork(self) -> "BlockTable":
"""Creates a new BlockTable instance with a copy of the blocks from the """Creates a new BlockTable instance with a copy of the blocks from the
current instance. current instance.
This method creates a new BlockTable instance with the same block size, This method creates a new BlockTable instance with the same block size,
block allocator, and a copy of the blocks from the current instance. The block allocator, and a copy of the blocks from the current instance. The
new BlockTable has its own independent set of blocks, but shares the new BlockTable has its own independent set of blocks, but shares the
same underlying memory allocation with the original BlockTable. same underlying memory allocation with the original BlockTable.
Returns: Returns:
BlockTable: A new BlockTable instance with a copy of the blocks from BlockTable: A new BlockTable instance with a copy of the blocks from
the current instance. the current instance.
""" """
assert self._is_allocated assert self._is_allocated
assert len(self._blocks) > 0 assert len(self._blocks) > 0
forked_blocks = self._allocator.fork(self._blocks[-1]) forked_blocks = self._allocator.fork(self._blocks[-1])
return BlockTable( return BlockTable(
block_size=self._block_size, block_size=self._block_size,
@@ -223,84 +223,84 @@ class BlockTable:
max_block_sliding_window=self._max_block_sliding_window, max_block_sliding_window=self._max_block_sliding_window,
cache_namespace=self._cache_namespace, cache_namespace=self._cache_namespace,
) )
def free(self) -> None: def free(self) -> None:
"""Frees the memory occupied by the blocks in the BlockTable. """Frees the memory occupied by the blocks in the BlockTable.
This method iterates over all the blocks in the `_blocks` list and calls This method iterates over all the blocks in the `_blocks` list and calls
the `free` method of the `_allocator` object to release the memory the `free` method of the `_allocator` object to release the memory
occupied by each block. After freeing all the blocks, the `_blocks` list occupied by each block. After freeing all the blocks, the `_blocks` list
is set to `None`. is set to `None`.
""" """
for block in self.blocks: for block in self.blocks:
self._allocator.free(block) self._allocator.free(block)
self._blocks.reset() self._blocks.reset()
@property @property
def physical_block_ids(self) -> List[int]: def physical_block_ids(self) -> List[int]:
"""Returns a list of physical block indices for the blocks in the """Returns a list of physical block indices for the blocks in the
BlockTable. BlockTable.
This property returns a list of integers, where each integer represents This property returns a list of integers, where each integer represents
the physical block index of a corresponding block in the `_blocks` list. the physical block index of a corresponding block in the `_blocks` list.
The physical block index is a unique identifier for the memory location The physical block index is a unique identifier for the memory location
occupied by the block. occupied by the block.
Returns: Returns:
List[int]: A list of physical block indices for the blocks in the List[int]: A list of physical block indices for the blocks in the
BlockTable. BlockTable.
""" """
return self._blocks.ids() return self._blocks.ids()
def get_unseen_token_ids(self, sequence_token_ids: List[int]) -> List[int]: def get_unseen_token_ids(self, sequence_token_ids: List[int]) -> List[int]:
"""Get the number of "unseen" tokens in the sequence. """Get the number of "unseen" tokens in the sequence.
Unseen tokens are tokens in the sequence corresponding to this block Unseen tokens are tokens in the sequence corresponding to this block
table, but are not yet appended to this block table. table, but are not yet appended to this block table.
Args: Args:
sequence_token_ids (List[int]): The list of token ids in the sequence_token_ids (List[int]): The list of token ids in the
sequence. sequence.
Returns: Returns:
List[int]: The postfix of sequence_token_ids that has not yet been List[int]: The postfix of sequence_token_ids that has not yet been
appended to the block table. appended to the block table.
""" """
# Since the block table is append-only, the unseen token ids are the # Since the block table is append-only, the unseen token ids are the
# ones after the appended ones. # ones after the appended ones.
return sequence_token_ids[self.num_full_slots:] return sequence_token_ids[self.num_full_slots:]
def _allocate_blocks_for_token_ids(self, prev_block: Optional[Block], def _allocate_blocks_for_token_ids(self, prev_block: Optional[Block],
token_ids: List[int], token_ids: List[int],
device: Device) -> List[Block]: device: Device) -> List[Block]:
blocks: List[Block] = [] blocks: List[Block] = []
block_token_ids = [] block_token_ids = []
tail_token_ids = [] tail_token_ids = []
for cur_token_ids in chunk_list(token_ids, self._block_size): for cur_token_ids in chunk_list(token_ids, self._block_size):
if len(cur_token_ids) == self._block_size: if len(cur_token_ids) == self._block_size:
block_token_ids.append(cur_token_ids) block_token_ids.append(cur_token_ids)
else: else:
tail_token_ids.append(cur_token_ids) tail_token_ids.append(cur_token_ids)
if block_token_ids: if block_token_ids:
blocks.extend(self._allocate_immutable_blocks( blocks.extend(self._allocate_immutable_blocks(
prev_block=prev_block, prev_block=prev_block,
block_token_ids=block_token_ids, block_token_ids=block_token_ids,
device=device)) device=device))
prev_block = blocks[-1] prev_block = blocks[-1]
if tail_token_ids: if tail_token_ids:
assert len(tail_token_ids) == 1 assert len(tail_token_ids) == 1
cur_token_ids = tail_token_ids[0] cur_token_ids = tail_token_ids[0]
block = self._allocate_mutable_block(prev_block=prev_block, block = self._allocate_mutable_block(prev_block=prev_block,
device=device) device=device)
block.append_token_ids(cur_token_ids) block.append_token_ids(cur_token_ids)
blocks.append(block) blocks.append(block)
return blocks return blocks
def _allocate_mutable_block(self, prev_block: Optional[Block], def _allocate_mutable_block(self, prev_block: Optional[Block],
@@ -372,85 +372,85 @@ class BlockTable:
prev_block, prev_block,
block_token_ids=block_token_ids, block_token_ids=block_token_ids,
device=device) device=device)
def _get_all_token_ids(self) -> List[int]: def _get_all_token_ids(self) -> List[int]:
# NOTE: This function is O(seq_len); use sparingly. # NOTE: This function is O(seq_len); use sparingly.
token_ids: List[int] = [] token_ids: List[int] = []
if not self._is_allocated: if not self._is_allocated:
return token_ids return token_ids
for block in self.blocks: for block in self.blocks:
token_ids.extend(block.token_ids) token_ids.extend(block.token_ids)
return token_ids return token_ids
def _get_num_token_ids(self) -> int: def _get_num_token_ids(self) -> int:
res = 0 res = 0
for block in self.blocks: for block in self.blocks:
res += len(block.token_ids) res += len(block.token_ids)
return res return res
@property @property
def _is_allocated(self) -> bool: def _is_allocated(self) -> bool:
return len(self._blocks) > 0 return len(self._blocks) > 0
@property @property
def blocks(self) -> List[Block]: def blocks(self) -> List[Block]:
return self._blocks.list() return self._blocks.list()
@property @property
def _num_empty_slots(self) -> int: def _num_empty_slots(self) -> int:
assert self._is_allocated assert self._is_allocated
return len(self._blocks) * self._block_size - self._num_full_slots return len(self._blocks) * self._block_size - self._num_full_slots
@property @property
def num_full_slots(self) -> int: def num_full_slots(self) -> int:
"""Returns the total number of tokens currently stored in the """Returns the total number of tokens currently stored in the
BlockTable. BlockTable.
Returns: Returns:
int: The total number of tokens currently stored in the BlockTable. int: The total number of tokens currently stored in the BlockTable.
""" """
return self._num_full_slots return self._num_full_slots
def get_num_blocks_touched_by_append_slots( def get_num_blocks_touched_by_append_slots(
self, token_ids: List[int], num_lookahead_slots: int) -> int: self, token_ids: List[int], num_lookahead_slots: int) -> int:
"""Determine how many blocks will be "touched" by appending the token """Determine how many blocks will be "touched" by appending the token
ids. ids.
This is required for the scheduler to determine whether a sequence can This is required for the scheduler to determine whether a sequence can
continue generation, or if it must be preempted. continue generation, or if it must be preempted.
""" """
# Math below is equivalent to: # Math below is equivalent to:
# all_token_ids = token_ids + [-1] * num_lookahead_slots # all_token_ids = token_ids + [-1] * num_lookahead_slots
# token_blocks = self._chunk_token_blocks_for_append(all_token_ids) # token_blocks = self._chunk_token_blocks_for_append(all_token_ids)
# return len(token_blocks) # return len(token_blocks)
num_token_ids = len(token_ids) + num_lookahead_slots num_token_ids = len(token_ids) + num_lookahead_slots
first_chunk_size = self._block_size - (self._num_full_slots % first_chunk_size = self._block_size - (self._num_full_slots %
self._block_size) self._block_size)
num_token_blocks = (1 + math.ceil( num_token_blocks = (1 + math.ceil(
(num_token_ids - first_chunk_size) / self._block_size)) (num_token_ids - first_chunk_size) / self._block_size))
return num_token_blocks return num_token_blocks
def _chunk_token_blocks_for_append( def _chunk_token_blocks_for_append(
self, token_ids: List[int]) -> List[List[int]]: self, token_ids: List[int]) -> List[List[int]]:
"""Split the token ids into block-sized chunks so they can be easily """Split the token ids into block-sized chunks so they can be easily
appended to blocks. The first such "token block" may have less token ids appended to blocks. The first such "token block" may have less token ids
than the block size, since the last allocated block may be partially than the block size, since the last allocated block may be partially
full. full.
If no token ids are provided, then no chunks are returned. If no token ids are provided, then no chunks are returned.
""" """
if not token_ids: if not token_ids:
return [] return []
first_chunk_size = self._block_size - (self._num_full_slots % first_chunk_size = self._block_size - (self._num_full_slots %
self._block_size) self._block_size)
token_blocks = [token_ids[:first_chunk_size]] token_blocks = [token_ids[:first_chunk_size]]
token_blocks.extend( token_blocks.extend(
chunk_list(token_ids[first_chunk_size:], self._block_size)) chunk_list(token_ids[first_chunk_size:], self._block_size))
return token_blocks return token_blocks

View File

@@ -4,56 +4,56 @@ from vllm.core.block.cpu_kv_content_cache import (CpuKvContentCache,
cpu_kv_offload_enabled) cpu_kv_offload_enabled)
from vllm.core.block.interfaces import (Block, BlockAllocator, BlockId, from vllm.core.block.interfaces import (Block, BlockAllocator, BlockId,
DeviceAwareBlockAllocator) DeviceAwareBlockAllocator)
from vllm.core.block.naive_block import NaiveBlock, NaiveBlockAllocator from vllm.core.block.naive_block import NaiveBlock, NaiveBlockAllocator
from vllm.core.block.prefix_caching_block import PrefixCachingBlockAllocator from vllm.core.block.prefix_caching_block import PrefixCachingBlockAllocator
from vllm.utils import Device from vllm.utils import Device
class CpuGpuBlockAllocator(DeviceAwareBlockAllocator): class CpuGpuBlockAllocator(DeviceAwareBlockAllocator):
"""A block allocator that can allocate blocks on both CPU and GPU memory. """A block allocator that can allocate blocks on both CPU and GPU memory.
This class implements the `DeviceAwareBlockAllocator` interface and provides This class implements the `DeviceAwareBlockAllocator` interface and provides
functionality for allocating and managing blocks of memory on both CPU and functionality for allocating and managing blocks of memory on both CPU and
GPU devices. GPU devices.
The `CpuGpuBlockAllocator` maintains separate memory pools for CPU and GPU The `CpuGpuBlockAllocator` maintains separate memory pools for CPU and GPU
blocks, and allows for allocation, deallocation, forking, and swapping of blocks, and allows for allocation, deallocation, forking, and swapping of
blocks across these memory pools. blocks across these memory pools.
""" """
@staticmethod @staticmethod
def create( def create(
allocator_type: str, allocator_type: str,
num_gpu_blocks: int, num_gpu_blocks: int,
num_cpu_blocks: int, num_cpu_blocks: int,
block_size: int, block_size: int,
) -> DeviceAwareBlockAllocator: ) -> DeviceAwareBlockAllocator:
"""Creates a CpuGpuBlockAllocator instance with the specified """Creates a CpuGpuBlockAllocator instance with the specified
configuration. configuration.
This static method creates and returns a CpuGpuBlockAllocator instance This static method creates and returns a CpuGpuBlockAllocator instance
based on the provided parameters. It initializes the CPU and GPU block based on the provided parameters. It initializes the CPU and GPU block
allocators with the specified number of blocks, block size, and allocators with the specified number of blocks, block size, and
allocator type. allocator type.
Args: Args:
allocator_type (str): The type of block allocator to use for CPU allocator_type (str): The type of block allocator to use for CPU
and GPU blocks. Currently supported values are "naive" and and GPU blocks. Currently supported values are "naive" and
"prefix_caching". "prefix_caching".
num_gpu_blocks (int): The number of blocks to allocate for GPU num_gpu_blocks (int): The number of blocks to allocate for GPU
memory. memory.
num_cpu_blocks (int): The number of blocks to allocate for CPU num_cpu_blocks (int): The number of blocks to allocate for CPU
memory. memory.
block_size (int): The size of each block in number of tokens. block_size (int): The size of each block in number of tokens.
Returns: Returns:
DeviceAwareBlockAllocator: A CpuGpuBlockAllocator instance with the DeviceAwareBlockAllocator: A CpuGpuBlockAllocator instance with the
specified configuration. specified configuration.
Notes: Notes:
- The block IDs are assigned contiguously, with GPU block IDs coming - The block IDs are assigned contiguously, with GPU block IDs coming
before CPU block IDs. before CPU block IDs.
""" """
content_offload = cpu_kv_offload_enabled() content_offload = cpu_kv_offload_enabled()
if content_offload and allocator_type != "prefix_caching": if content_offload and allocator_type != "prefix_caching":
raise RuntimeError( raise RuntimeError(
@@ -63,38 +63,38 @@ class CpuGpuBlockAllocator(DeviceAwareBlockAllocator):
"BI100_CPU_KV_OFFLOAD=1 requires at least one CPU KV block") "BI100_CPU_KV_OFFLOAD=1 requires at least one CPU KV block")
block_ids = list(range(num_gpu_blocks + num_cpu_blocks)) block_ids = list(range(num_gpu_blocks + num_cpu_blocks))
gpu_block_ids = block_ids[:num_gpu_blocks] gpu_block_ids = block_ids[:num_gpu_blocks]
cpu_block_ids = block_ids[num_gpu_blocks:] cpu_block_ids = block_ids[num_gpu_blocks:]
if allocator_type == "naive": if allocator_type == "naive":
gpu_allocator: BlockAllocator = NaiveBlockAllocator( gpu_allocator: BlockAllocator = NaiveBlockAllocator(
create_block=NaiveBlock, # type: ignore create_block=NaiveBlock, # type: ignore
num_blocks=num_gpu_blocks, num_blocks=num_gpu_blocks,
block_size=block_size, block_size=block_size,
block_ids=gpu_block_ids, block_ids=gpu_block_ids,
) )
cpu_allocator: BlockAllocator = NaiveBlockAllocator( cpu_allocator: BlockAllocator = NaiveBlockAllocator(
create_block=NaiveBlock, # type: ignore create_block=NaiveBlock, # type: ignore
num_blocks=num_cpu_blocks, num_blocks=num_cpu_blocks,
block_size=block_size, block_size=block_size,
block_ids=cpu_block_ids, block_ids=cpu_block_ids,
) )
elif allocator_type == "prefix_caching": elif allocator_type == "prefix_caching":
gpu_allocator = PrefixCachingBlockAllocator( gpu_allocator = PrefixCachingBlockAllocator(
num_blocks=num_gpu_blocks, num_blocks=num_gpu_blocks,
block_size=block_size, block_size=block_size,
block_ids=gpu_block_ids, block_ids=gpu_block_ids,
) )
cpu_allocator = PrefixCachingBlockAllocator( cpu_allocator = PrefixCachingBlockAllocator(
num_blocks=num_cpu_blocks, num_blocks=num_cpu_blocks,
block_size=block_size, block_size=block_size,
block_ids=cpu_block_ids, block_ids=cpu_block_ids,
) )
else: else:
raise ValueError(f"Unknown allocator type {allocator_type=}") raise ValueError(f"Unknown allocator type {allocator_type=}")
return CpuGpuBlockAllocator( return CpuGpuBlockAllocator(
cpu_block_allocator=cpu_allocator, cpu_block_allocator=cpu_allocator,
gpu_block_allocator=gpu_allocator, gpu_block_allocator=gpu_allocator,
@@ -105,21 +105,21 @@ class CpuGpuBlockAllocator(DeviceAwareBlockAllocator):
def __init__(self, cpu_block_allocator: BlockAllocator, def __init__(self, cpu_block_allocator: BlockAllocator,
gpu_block_allocator: BlockAllocator, gpu_block_allocator: BlockAllocator,
cpu_content_cache: Optional[CpuKvContentCache] = None): cpu_content_cache: Optional[CpuKvContentCache] = None):
assert not ( assert not (
cpu_block_allocator.all_block_ids cpu_block_allocator.all_block_ids
& gpu_block_allocator.all_block_ids & gpu_block_allocator.all_block_ids
), "cpu and gpu block allocators can't have intersection of block ids" ), "cpu and gpu block allocators can't have intersection of block ids"
self._allocators = { self._allocators = {
Device.CPU: cpu_block_allocator, Device.CPU: cpu_block_allocator,
Device.GPU: gpu_block_allocator, Device.GPU: gpu_block_allocator,
} }
self._swap_mapping: Dict[int, int] = {} self._swap_mapping: Dict[int, int] = {}
self._null_block: Optional[Block] = None self._null_block: Optional[Block] = None
self._cpu_content_cache = cpu_content_cache self._cpu_content_cache = cpu_content_cache
self._block_ids_to_allocator: Dict[int, BlockAllocator] = {} self._block_ids_to_allocator: Dict[int, BlockAllocator] = {}
for _, allocator in self._allocators.items(): for _, allocator in self._allocators.items():
for block_id in allocator.all_block_ids: for block_id in allocator.all_block_ids:
self._block_ids_to_allocator[block_id] = allocator self._block_ids_to_allocator[block_id] = allocator
@@ -164,236 +164,236 @@ class CpuGpuBlockAllocator(DeviceAwareBlockAllocator):
assert self._cpu_content_cache is not None assert self._cpu_content_cache is not None
gpu_slot = self.get_physical_block_id(Device.GPU, gpu_block_id) gpu_slot = self.get_physical_block_id(Device.GPU, gpu_block_id)
return self._cpu_content_cache.stage_store(content_hash, gpu_slot) return self._cpu_content_cache.stage_store(content_hash, gpu_slot)
def allocate_or_get_null_block(self) -> Block: def allocate_or_get_null_block(self) -> Block:
if self._null_block is None: if self._null_block is None:
self._null_block = NullBlock( self._null_block = NullBlock(
self.allocate_mutable_block(None, Device.GPU)) self.allocate_mutable_block(None, Device.GPU))
return self._null_block return self._null_block
def allocate_mutable_block(self, prev_block: Optional[Block], def allocate_mutable_block(self, prev_block: Optional[Block],
device: Device) -> Block: device: Device) -> Block:
"""Allocates a new mutable block on the specified device. """Allocates a new mutable block on the specified device.
Args: Args:
prev_block (Optional[Block]): The previous block to in the sequence. prev_block (Optional[Block]): The previous block to in the sequence.
Used for prefix hashing. Used for prefix hashing.
device (Device): The device on which to allocate the new block. device (Device): The device on which to allocate the new block.
Returns: Returns:
Block: The newly allocated mutable block. Block: The newly allocated mutable block.
""" """
return self._allocators[device].allocate_mutable_block(prev_block) return self._allocators[device].allocate_mutable_block(prev_block)
def allocate_immutable_blocks(self, prev_block: Optional[Block], def allocate_immutable_blocks(self, prev_block: Optional[Block],
block_token_ids: List[List[int]], block_token_ids: List[List[int]],
device: Device) -> List[Block]: device: Device) -> List[Block]:
"""Allocates a new group of immutable blocks with the provided block """Allocates a new group of immutable blocks with the provided block
token IDs on the specified device. token IDs on the specified device.
Args: Args:
prev_block (Optional[Block]): The previous block in the sequence. prev_block (Optional[Block]): The previous block in the sequence.
Used for prefix hashing. Used for prefix hashing.
block_token_ids (List[int]): The list of block token IDs to be block_token_ids (List[int]): The list of block token IDs to be
stored in the new blocks. stored in the new blocks.
device (Device): The device on which to allocate the new block. device (Device): The device on which to allocate the new block.
Returns: Returns:
List[Block]: The newly allocated list of immutable blocks List[Block]: The newly allocated list of immutable blocks
containing the provided block token IDs. containing the provided block token IDs.
""" """
return self._allocators[device].allocate_immutable_blocks( return self._allocators[device].allocate_immutable_blocks(
prev_block, block_token_ids) prev_block, block_token_ids)
def allocate_immutable_block(self, prev_block: Optional[Block], def allocate_immutable_block(self, prev_block: Optional[Block],
token_ids: List[int], token_ids: List[int],
device: Device) -> Block: device: Device) -> Block:
"""Allocates a new immutable block with the provided token IDs on the """Allocates a new immutable block with the provided token IDs on the
specified device. specified device.
Args: Args:
prev_block (Optional[Block]): The previous block in the sequence. prev_block (Optional[Block]): The previous block in the sequence.
Used for prefix hashing. Used for prefix hashing.
token_ids (List[int]): The list of token IDs to be stored in the new token_ids (List[int]): The list of token IDs to be stored in the new
block. block.
device (Device): The device on which to allocate the new block. device (Device): The device on which to allocate the new block.
Returns: Returns:
Block: The newly allocated immutable block containing the provided Block: The newly allocated immutable block containing the provided
token IDs. token IDs.
""" """
return self._allocators[device].allocate_immutable_block( return self._allocators[device].allocate_immutable_block(
prev_block, token_ids) prev_block, token_ids)
def free(self, block: Block) -> None: def free(self, block: Block) -> None:
"""Frees the memory occupied by the given block. """Frees the memory occupied by the given block.
Args: Args:
block (Block): The block to be freed. block (Block): The block to be freed.
""" """
# Null block should never be freed # Null block should never be freed
if isinstance(block, NullBlock): if isinstance(block, NullBlock):
return return
block_id = block.block_id block_id = block.block_id
assert block_id is not None assert block_id is not None
allocator = self._block_ids_to_allocator[block_id] allocator = self._block_ids_to_allocator[block_id]
allocator.free(block) allocator.free(block)
def fork(self, last_block: Block) -> List[Block]: def fork(self, last_block: Block) -> List[Block]:
"""Creates a new sequence of blocks that shares the same underlying """Creates a new sequence of blocks that shares the same underlying
memory as the original sequence. memory as the original sequence.
Args: Args:
last_block (Block): The last block in the original sequence. last_block (Block): The last block in the original sequence.
Returns: Returns:
List[Block]: A new list of blocks that shares the same memory as the List[Block]: A new list of blocks that shares the same memory as the
original sequence. original sequence.
""" """
# do not attempt to fork the null block # do not attempt to fork the null block
assert not isinstance(last_block, NullBlock) assert not isinstance(last_block, NullBlock)
block_id = last_block.block_id block_id = last_block.block_id
assert block_id is not None assert block_id is not None
allocator = self._block_ids_to_allocator[block_id] allocator = self._block_ids_to_allocator[block_id]
return allocator.fork(last_block) return allocator.fork(last_block)
def get_num_free_blocks(self, device: Device) -> int: def get_num_free_blocks(self, device: Device) -> int:
"""Returns the number of free blocks available on the specified device. """Returns the number of free blocks available on the specified device.
Args: Args:
device (Device): The device for which to query the number of free device (Device): The device for which to query the number of free
blocks. AssertionError is raised if None is passed. blocks. AssertionError is raised if None is passed.
Returns: Returns:
int: The number of free blocks available on the specified device. int: The number of free blocks available on the specified device.
""" """
return self._allocators[device].get_num_free_blocks() return self._allocators[device].get_num_free_blocks()
def get_num_total_blocks(self, device: Device) -> int: def get_num_total_blocks(self, device: Device) -> int:
return self._allocators[device].get_num_total_blocks() return self._allocators[device].get_num_total_blocks()
def get_physical_block_id(self, device: Device, absolute_id: int) -> int: def get_physical_block_id(self, device: Device, absolute_id: int) -> int:
"""Returns the zero-offset block id on certain device given the """Returns the zero-offset block id on certain device given the
absolute block id. absolute block id.
Args: Args:
device (Device): The device for which to query relative block id. device (Device): The device for which to query relative block id.
absolute_id (int): The absolute block id for the block in absolute_id (int): The absolute block id for the block in
whole allocator. whole allocator.
Returns: Returns:
int: The zero-offset block id on certain device. int: The zero-offset block id on certain device.
""" """
return self._allocators[device].get_physical_block_id(absolute_id) return self._allocators[device].get_physical_block_id(absolute_id)
def swap(self, blocks: List[Block], src_device: Device, def swap(self, blocks: List[Block], src_device: Device,
dst_device: Device) -> Dict[int, int]: dst_device: Device) -> Dict[int, int]:
"""Execute the swap for the given blocks from source_device """Execute the swap for the given blocks from source_device
on to dest_device, save the current swap mapping and append on to dest_device, save the current swap mapping and append
them to the accumulated `self._swap_mapping` for each them to the accumulated `self._swap_mapping` for each
scheduling move. scheduling move.
Args: Args:
blocks: List of blocks to be swapped. blocks: List of blocks to be swapped.
src_device (Device): Device to swap the 'blocks' from. src_device (Device): Device to swap the 'blocks' from.
dst_device (Device): Device to swap the 'blocks' to. dst_device (Device): Device to swap the 'blocks' to.
Returns: Returns:
Dict[int, int]: Swap mapping from source_device Dict[int, int]: Swap mapping from source_device
on to dest_device. on to dest_device.
""" """
if self.content_offload_enabled: if self.content_offload_enabled:
raise RuntimeError( raise RuntimeError(
"request-level preemption swap cannot share CPU slots with " "request-level preemption swap cannot share CPU slots with "
"BI100_CPU_KV_OFFLOAD") "BI100_CPU_KV_OFFLOAD")
src_block_ids = [block.block_id for block in blocks] src_block_ids = [block.block_id for block in blocks]
self._allocators[src_device].swap_out(blocks) self._allocators[src_device].swap_out(blocks)
self._allocators[dst_device].swap_in(blocks) self._allocators[dst_device].swap_in(blocks)
dst_block_ids = [block.block_id for block in blocks] dst_block_ids = [block.block_id for block in blocks]
current_swap_mapping: Dict[int, int] = {} current_swap_mapping: Dict[int, int] = {}
for src_block_id, dst_block_id in zip(src_block_ids, dst_block_ids): for src_block_id, dst_block_id in zip(src_block_ids, dst_block_ids):
if src_block_id is not None and dst_block_id is not None: if src_block_id is not None and dst_block_id is not None:
self._swap_mapping[src_block_id] = dst_block_id self._swap_mapping[src_block_id] = dst_block_id
current_swap_mapping[src_block_id] = dst_block_id current_swap_mapping[src_block_id] = dst_block_id
return current_swap_mapping return current_swap_mapping
def get_num_full_blocks_touched(self, blocks: List[Block], def get_num_full_blocks_touched(self, blocks: List[Block],
device: Device) -> int: device: Device) -> int:
"""Returns the number of full blocks that will be touched by """Returns the number of full blocks that will be touched by
swapping in/out the given blocks on to the 'device'. swapping in/out the given blocks on to the 'device'.
Args: Args:
blocks: List of blocks to be swapped. blocks: List of blocks to be swapped.
device (Device): Device to swap the 'blocks' on. device (Device): Device to swap the 'blocks' on.
Returns: Returns:
int: the number of full blocks that will be touched by int: the number of full blocks that will be touched by
swapping in/out the given blocks on to the 'device'. swapping in/out the given blocks on to the 'device'.
Non full blocks are ignored when deciding the number Non full blocks are ignored when deciding the number
of blocks to touch. of blocks to touch.
""" """
return self._allocators[device].get_num_full_blocks_touched(blocks) return self._allocators[device].get_num_full_blocks_touched(blocks)
def clear_copy_on_writes(self) -> List[Tuple[int, int]]: def clear_copy_on_writes(self) -> List[Tuple[int, int]]:
"""Clears the copy-on-write (CoW) state and returns the mapping of """Clears the copy-on-write (CoW) state and returns the mapping of
source to destination block IDs. source to destination block IDs.
Returns: Returns:
List[Tuple[int, int]]: A list mapping source block IDs to List[Tuple[int, int]]: A list mapping source block IDs to
destination block IDs. destination block IDs.
""" """
# CoW only supported on GPU # CoW only supported on GPU
device = Device.GPU device = Device.GPU
return self._allocators[device].clear_copy_on_writes() return self._allocators[device].clear_copy_on_writes()
def mark_blocks_as_accessed(self, block_ids: List[int], def mark_blocks_as_accessed(self, block_ids: List[int],
now: float) -> None: now: float) -> None:
"""Mark blocks as accessed, only use for prefix caching.""" """Mark blocks as accessed, only use for prefix caching."""
# Prefix caching only supported on GPU. # Prefix caching only supported on GPU.
device = Device.GPU device = Device.GPU
return self._allocators[device].mark_blocks_as_accessed(block_ids, now) return self._allocators[device].mark_blocks_as_accessed(block_ids, now)
def mark_blocks_as_computed(self, block_ids: List[int]) -> None: def mark_blocks_as_computed(self, block_ids: List[int]) -> None:
"""Mark blocks as accessed, only use for prefix caching.""" """Mark blocks as accessed, only use for prefix caching."""
# Prefix caching only supported on GPU. # Prefix caching only supported on GPU.
device = Device.GPU device = Device.GPU
return self._allocators[device].mark_blocks_as_computed(block_ids) return self._allocators[device].mark_blocks_as_computed(block_ids)
def get_computed_block_ids(self, prev_computed_block_ids: List[int], def get_computed_block_ids(self, prev_computed_block_ids: List[int],
block_ids: List[int], block_ids: List[int],
skip_last_block_id: bool) -> List[int]: skip_last_block_id: bool) -> List[int]:
# Prefix caching only supported on GPU. # Prefix caching only supported on GPU.
device = Device.GPU device = Device.GPU
return self._allocators[device].get_computed_block_ids( return self._allocators[device].get_computed_block_ids(
prev_computed_block_ids, block_ids, skip_last_block_id) prev_computed_block_ids, block_ids, skip_last_block_id)
def get_common_computed_block_ids( def get_common_computed_block_ids(
self, computed_seq_block_ids: List[List[int]]) -> List[int]: self, computed_seq_block_ids: List[List[int]]) -> List[int]:
# Prefix caching only supported on GPU. # Prefix caching only supported on GPU.
device = Device.GPU device = Device.GPU
return self._allocators[device].get_common_computed_block_ids( return self._allocators[device].get_common_computed_block_ids(
computed_seq_block_ids) computed_seq_block_ids)
@property @property
def all_block_ids(self) -> FrozenSet[int]: def all_block_ids(self) -> FrozenSet[int]:
return frozenset(self._block_ids_to_allocator.keys()) return frozenset(self._block_ids_to_allocator.keys())
def get_prefix_cache_hit_rate(self, device: Device) -> float: def get_prefix_cache_hit_rate(self, device: Device) -> float:
"""Prefix cache hit rate. -1 means not supported or disabled.""" """Prefix cache hit rate. -1 means not supported or disabled."""
assert device in self._allocators assert device in self._allocators
return self._allocators[device].get_prefix_cache_hit_rate() return self._allocators[device].get_prefix_cache_hit_rate()
def get_and_reset_swaps(self) -> List[Tuple[int, int]]: def get_and_reset_swaps(self) -> List[Tuple[int, int]]:
"""Returns and clears the mapping of source to destination block IDs. """Returns and clears the mapping of source to destination block IDs.
Will be called after every swapping operations for now, and after every Will be called after every swapping operations for now, and after every
schedule when BlockManagerV2 become default. Currently not useful. schedule when BlockManagerV2 become default. Currently not useful.
Returns: Returns:
List[Tuple[int, int]]: A mapping of source to destination block IDs. List[Tuple[int, int]]: A mapping of source to destination block IDs.
""" """
mapping = self._swap_mapping.copy() mapping = self._swap_mapping.copy()
self._swap_mapping.clear() self._swap_mapping.clear()
return list(mapping.items()) return list(mapping.items())
@@ -407,69 +407,69 @@ class CpuGpuBlockAllocator(DeviceAwareBlockAllocator):
def begin_prefix_cache_step(self) -> None: def begin_prefix_cache_step(self) -> None:
if self._cpu_content_cache is not None: if self._cpu_content_cache is not None:
self._cpu_content_cache.begin_step() self._cpu_content_cache.begin_step()
class NullBlock(Block): class NullBlock(Block):
""" """
Null blocks are used as a placeholders for KV cache blocks that have Null blocks are used as a placeholders for KV cache blocks that have
been dropped due to sliding window. been dropped due to sliding window.
This implementation just wraps an ordinary block and prevents it from This implementation just wraps an ordinary block and prevents it from
being modified. It also allows for testing if a block is NullBlock being modified. It also allows for testing if a block is NullBlock
via isinstance(). via isinstance().
""" """
def __init__(self, proxy: Block): def __init__(self, proxy: Block):
super().__init__() super().__init__()
self._proxy = proxy self._proxy = proxy
def append_token_ids(self, token_ids: List[BlockId]): def append_token_ids(self, token_ids: List[BlockId]):
raise ValueError("null block should not be modified") raise ValueError("null block should not be modified")
@property @property
def block_id(self): def block_id(self):
return self._proxy.block_id return self._proxy.block_id
@block_id.setter @block_id.setter
def block_id(self, value: Optional[BlockId]): def block_id(self, value: Optional[BlockId]):
raise ValueError("null block should not be modified") raise ValueError("null block should not be modified")
@property @property
def token_ids(self) -> List[BlockId]: def token_ids(self) -> List[BlockId]:
return self._proxy.token_ids return self._proxy.token_ids
@property @property
def num_tokens_total(self) -> int: def num_tokens_total(self) -> int:
raise NotImplementedError( raise NotImplementedError(
"num_tokens_total is not used for null block") "num_tokens_total is not used for null block")
@property @property
def num_empty_slots(self) -> BlockId: def num_empty_slots(self) -> BlockId:
return self._proxy.num_empty_slots return self._proxy.num_empty_slots
@property @property
def is_full(self): def is_full(self):
return self._proxy.is_full return self._proxy.is_full
@property @property
def prev_block(self): def prev_block(self):
return self._proxy.prev_block return self._proxy.prev_block
@property @property
def computed(self): def computed(self):
return self._proxy.computed return self._proxy.computed
@computed.setter @computed.setter
def computed(self, value): def computed(self, value):
self._proxy.computed = value self._proxy.computed = value
@property @property
def last_accessed(self) -> float: def last_accessed(self) -> float:
return self._proxy.last_accessed return self._proxy.last_accessed
@last_accessed.setter @last_accessed.setter
def last_accessed(self, last_accessed_ts: float): def last_accessed(self, last_accessed_ts: float):
self._proxy.last_accessed = last_accessed_ts self._proxy.last_accessed = last_accessed_ts
@property @property
def content_hash(self): def content_hash(self):
return self._proxy.content_hash return self._proxy.content_hash

View File

@@ -32,86 +32,86 @@ logger = init_logger(__name__)
class BlockSpaceManagerV2(BlockSpaceManager): class BlockSpaceManagerV2(BlockSpaceManager):
"""BlockSpaceManager which manages the allocation of KV cache. """BlockSpaceManager which manages the allocation of KV cache.
It owns responsibility for allocation, swapping, allocating memory for It owns responsibility for allocation, swapping, allocating memory for
autoregressively-generated tokens, and other advanced features such as autoregressively-generated tokens, and other advanced features such as
prefix caching, forking/copy-on-write, and sliding-window memory allocation. prefix caching, forking/copy-on-write, and sliding-window memory allocation.
This class implements the design described in This class implements the design described in
https://github.com/vllm-project/vllm/pull/3492. https://github.com/vllm-project/vllm/pull/3492.
Lookahead slots Lookahead slots
The block manager has the notion of a "lookahead slot". These are slots The block manager has the notion of a "lookahead slot". These are slots
in the KV cache that are allocated for a sequence. Unlike the other in the KV cache that are allocated for a sequence. Unlike the other
allocated slots, the content of these slots is undefined -- the worker allocated slots, the content of these slots is undefined -- the worker
may use the memory allocations in any way. may use the memory allocations in any way.
In practice, a worker could use these lookahead slots to run multiple In practice, a worker could use these lookahead slots to run multiple
forward passes for a single scheduler invocation. Each successive forward passes for a single scheduler invocation. Each successive
forward pass would write KV activations to the corresponding lookahead forward pass would write KV activations to the corresponding lookahead
slot. This allows low inter-token latency use-cases, where the overhead slot. This allows low inter-token latency use-cases, where the overhead
of continuous batching scheduling is amortized over >1 generated tokens. of continuous batching scheduling is amortized over >1 generated tokens.
Speculative decoding uses lookahead slots to store KV activations of Speculative decoding uses lookahead slots to store KV activations of
proposal tokens. proposal tokens.
See https://github.com/vllm-project/vllm/pull/3250 for more information See https://github.com/vllm-project/vllm/pull/3250 for more information
on lookahead scheduling. on lookahead scheduling.
Args: Args:
block_size (int): The size of each memory block. block_size (int): The size of each memory block.
num_gpu_blocks (int): The number of memory blocks allocated on GPU. num_gpu_blocks (int): The number of memory blocks allocated on GPU.
num_cpu_blocks (int): The number of memory blocks allocated on CPU. num_cpu_blocks (int): The number of memory blocks allocated on CPU.
watermark (float, optional): The threshold used for memory swapping. watermark (float, optional): The threshold used for memory swapping.
Defaults to 0.01. Defaults to 0.01.
sliding_window (Optional[int], optional): The size of the sliding sliding_window (Optional[int], optional): The size of the sliding
window. Defaults to None. window. Defaults to None.
enable_caching (bool, optional): Flag indicating whether caching is enable_caching (bool, optional): Flag indicating whether caching is
enabled. Defaults to False. enabled. Defaults to False.
""" """
def __init__( def __init__(
self, self,
block_size: int, block_size: int,
num_gpu_blocks: int, num_gpu_blocks: int,
num_cpu_blocks: int, num_cpu_blocks: int,
watermark: float = 0.01, watermark: float = 0.01,
sliding_window: Optional[int] = None, sliding_window: Optional[int] = None,
enable_caching: bool = False, enable_caching: bool = False,
) -> None: ) -> None:
self.block_size = block_size self.block_size = block_size
self.num_total_gpu_blocks = num_gpu_blocks self.num_total_gpu_blocks = num_gpu_blocks
self.num_total_cpu_blocks = num_cpu_blocks self.num_total_cpu_blocks = num_cpu_blocks
self.sliding_window = sliding_window self.sliding_window = sliding_window
# max_block_sliding_window is the max number of blocks that need to be # max_block_sliding_window is the max number of blocks that need to be
# allocated # allocated
self.max_block_sliding_window = None self.max_block_sliding_window = None
if sliding_window is not None: if sliding_window is not None:
# +1 here because // rounds down # +1 here because // rounds down
num_blocks = sliding_window // block_size + 1 num_blocks = sliding_window // block_size + 1
# +1 here because the last block may not be full, # +1 here because the last block may not be full,
# and so the sequence stretches one more block at the beginning # and so the sequence stretches one more block at the beginning
# For example, if sliding_window is 3 and block_size is 4, # For example, if sliding_window is 3 and block_size is 4,
# we may need 2 blocks when the second block only holds 1 token. # we may need 2 blocks when the second block only holds 1 token.
self.max_block_sliding_window = num_blocks + 1 self.max_block_sliding_window = num_blocks + 1
self.watermark = watermark self.watermark = watermark
assert watermark >= 0.0 assert watermark >= 0.0
self.enable_caching = enable_caching self.enable_caching = enable_caching
self.watermark_blocks = int(watermark * num_gpu_blocks) self.watermark_blocks = int(watermark * num_gpu_blocks)
self.block_allocator = CpuGpuBlockAllocator.create( self.block_allocator = CpuGpuBlockAllocator.create(
allocator_type="prefix_caching" if enable_caching else "naive", allocator_type="prefix_caching" if enable_caching else "naive",
num_gpu_blocks=num_gpu_blocks, num_gpu_blocks=num_gpu_blocks,
num_cpu_blocks=num_cpu_blocks, num_cpu_blocks=num_cpu_blocks,
block_size=block_size, block_size=block_size,
) )
self.block_tables: Dict[SeqId, BlockTable] = {} self.block_tables: Dict[SeqId, BlockTable] = {}
self.cross_block_tables: Dict[EncoderSeqId, BlockTable] = {} self.cross_block_tables: Dict[EncoderSeqId, BlockTable] = {}
self._warned_mm_namespace_requests = set[str]() self._warned_mm_namespace_requests = set[str]()
self._request_local_namespace: Dict[str, bytes] = {} self._request_local_namespace: Dict[str, bytes] = {}
@@ -121,46 +121,46 @@ class BlockSpaceManagerV2(BlockSpaceManager):
self.block_allocator) self.block_allocator)
self._last_access_blocks_tracker = LastAccessBlocksTracker( self._last_access_blocks_tracker = LastAccessBlocksTracker(
self.block_allocator) self.block_allocator)
def can_allocate(self, def can_allocate(self,
seq_group: SequenceGroup, seq_group: SequenceGroup,
num_lookahead_slots: int = 0) -> AllocStatus: num_lookahead_slots: int = 0) -> AllocStatus:
# FIXME(woosuk): Here we assume that all sequences in the group share # FIXME(woosuk): Here we assume that all sequences in the group share
# the same prompt. This may not be true for preempted sequences. # the same prompt. This may not be true for preempted sequences.
check_no_caching_or_swa_for_blockmgr_encdec(self, seq_group) check_no_caching_or_swa_for_blockmgr_encdec(self, seq_group)
seq = seq_group.get_seqs(status=SequenceStatus.WAITING)[0] seq = seq_group.get_seqs(status=SequenceStatus.WAITING)[0]
num_required_blocks = BlockTable.get_num_required_blocks( num_required_blocks = BlockTable.get_num_required_blocks(
seq.get_token_ids(), seq.get_token_ids(),
block_size=self.block_size, block_size=self.block_size,
num_lookahead_slots=num_lookahead_slots, num_lookahead_slots=num_lookahead_slots,
) )
if seq_group.is_encoder_decoder(): if seq_group.is_encoder_decoder():
encoder_seq = seq_group.get_encoder_seq() encoder_seq = seq_group.get_encoder_seq()
assert encoder_seq is not None assert encoder_seq is not None
num_required_blocks += BlockTable.get_num_required_blocks( num_required_blocks += BlockTable.get_num_required_blocks(
encoder_seq.get_token_ids(), encoder_seq.get_token_ids(),
block_size=self.block_size, block_size=self.block_size,
) )
if self.max_block_sliding_window is not None: if self.max_block_sliding_window is not None:
num_required_blocks = min(num_required_blocks, num_required_blocks = min(num_required_blocks,
self.max_block_sliding_window) self.max_block_sliding_window)
num_free_gpu_blocks = self.block_allocator.get_num_free_blocks( num_free_gpu_blocks = self.block_allocator.get_num_free_blocks(
device=Device.GPU) device=Device.GPU)
# Use watermark to avoid frequent cache eviction. # Use watermark to avoid frequent cache eviction.
if (self.num_total_gpu_blocks - num_required_blocks < if (self.num_total_gpu_blocks - num_required_blocks <
self.watermark_blocks): self.watermark_blocks):
return AllocStatus.NEVER return AllocStatus.NEVER
if num_free_gpu_blocks - num_required_blocks >= self.watermark_blocks: if num_free_gpu_blocks - num_required_blocks >= self.watermark_blocks:
return AllocStatus.OK return AllocStatus.OK
else: else:
return AllocStatus.LATER return AllocStatus.LATER
def _allocate_sequence( def _allocate_sequence(
self, self,
seq: Sequence, seq: Sequence,
@@ -181,10 +181,10 @@ class BlockSpaceManagerV2(BlockSpaceManager):
def allocate(self, seq_group: SequenceGroup) -> None: def allocate(self, seq_group: SequenceGroup) -> None:
# Allocate self-attention block tables for decoder sequences # Allocate self-attention block tables for decoder sequences
waiting_seqs = seq_group.get_seqs(status=SequenceStatus.WAITING) waiting_seqs = seq_group.get_seqs(status=SequenceStatus.WAITING)
assert not (set(seq.seq_id for seq in waiting_seqs) assert not (set(seq.seq_id for seq in waiting_seqs)
& self.block_tables.keys()), "block table already exists" & self.block_tables.keys()), "block table already exists"
# NOTE: Here we assume that all sequences in the group have the same # NOTE: Here we assume that all sequences in the group have the same
# prompt. # prompt.
seq = waiting_seqs[0] seq = waiting_seqs[0]
@@ -199,31 +199,31 @@ class BlockSpaceManagerV2(BlockSpaceManager):
cache_namespace=cache_namespace, cache_namespace=cache_namespace,
) )
self.block_tables[seq.seq_id] = block_table self.block_tables[seq.seq_id] = block_table
# Track seq # Track seq
self._computed_blocks_tracker.add_seq(seq.seq_id) self._computed_blocks_tracker.add_seq(seq.seq_id)
self._last_access_blocks_tracker.add_seq(seq.seq_id) self._last_access_blocks_tracker.add_seq(seq.seq_id)
# Assign the block table for each sequence. # Assign the block table for each sequence.
for seq in waiting_seqs[1:]: for seq in waiting_seqs[1:]:
self.block_tables[seq.seq_id] = block_table.fork() self.block_tables[seq.seq_id] = block_table.fork()
# Track seq # Track seq
self._computed_blocks_tracker.add_seq(seq.seq_id) self._computed_blocks_tracker.add_seq(seq.seq_id)
self._last_access_blocks_tracker.add_seq(seq.seq_id) self._last_access_blocks_tracker.add_seq(seq.seq_id)
# Allocate cross-attention block table for encoder sequence # Allocate cross-attention block table for encoder sequence
# #
# NOTE: Here we assume that all sequences in the group have the same # NOTE: Here we assume that all sequences in the group have the same
# encoder prompt. # encoder prompt.
request_id = seq_group.request_id request_id = seq_group.request_id
assert (request_id assert (request_id
not in self.cross_block_tables), \ not in self.cross_block_tables), \
"block table already exists" "block table already exists"
check_no_caching_or_swa_for_blockmgr_encdec(self, seq_group) check_no_caching_or_swa_for_blockmgr_encdec(self, seq_group)
if seq_group.is_encoder_decoder(): if seq_group.is_encoder_decoder():
encoder_seq = seq_group.get_encoder_seq() encoder_seq = seq_group.get_encoder_seq()
assert encoder_seq is not None assert encoder_seq is not None
@@ -279,129 +279,129 @@ class BlockSpaceManagerV2(BlockSpaceManager):
else: else:
digest.update(b"text|") digest.update(b"text|")
return digest.digest() return digest.digest()
def can_append_slots(self, seq_group: SequenceGroup, def can_append_slots(self, seq_group: SequenceGroup,
num_lookahead_slots: int) -> bool: num_lookahead_slots: int) -> bool:
"""Determine if there is enough space in the GPU KV cache to continue """Determine if there is enough space in the GPU KV cache to continue
generation of the specified sequence group. generation of the specified sequence group.
We use a worst-case heuristic: assume each touched block will require a We use a worst-case heuristic: assume each touched block will require a
new allocation (either via CoW or new block). We can append slots if the new allocation (either via CoW or new block). We can append slots if the
number of touched blocks is less than the number of free blocks. number of touched blocks is less than the number of free blocks.
"Lookahead slots" are slots that are allocated in addition to the slots "Lookahead slots" are slots that are allocated in addition to the slots
for known tokens. The contents of the lookahead slots are not defined. for known tokens. The contents of the lookahead slots are not defined.
This is used by speculative decoding when speculating future tokens. This is used by speculative decoding when speculating future tokens.
""" """
num_touched_blocks = 0 num_touched_blocks = 0
for seq in seq_group.get_seqs(status=SequenceStatus.RUNNING): for seq in seq_group.get_seqs(status=SequenceStatus.RUNNING):
block_table = self.block_tables[seq.seq_id] block_table = self.block_tables[seq.seq_id]
num_touched_blocks += ( num_touched_blocks += (
block_table.get_num_blocks_touched_by_append_slots( block_table.get_num_blocks_touched_by_append_slots(
token_ids=block_table.get_unseen_token_ids( token_ids=block_table.get_unseen_token_ids(
seq.get_token_ids()), seq.get_token_ids()),
num_lookahead_slots=num_lookahead_slots, num_lookahead_slots=num_lookahead_slots,
)) ))
num_free_gpu_blocks = self.block_allocator.get_num_free_blocks( num_free_gpu_blocks = self.block_allocator.get_num_free_blocks(
Device.GPU) Device.GPU)
return num_touched_blocks <= num_free_gpu_blocks return num_touched_blocks <= num_free_gpu_blocks
def append_slots( def append_slots(
self, self,
seq: Sequence, seq: Sequence,
num_lookahead_slots: int, num_lookahead_slots: int,
) -> List[Tuple[int, int]]: ) -> List[Tuple[int, int]]:
block_table = self.block_tables[seq.seq_id] block_table = self.block_tables[seq.seq_id]
block_table.append_token_ids( block_table.append_token_ids(
token_ids=block_table.get_unseen_token_ids(seq.get_token_ids()), token_ids=block_table.get_unseen_token_ids(seq.get_token_ids()),
num_lookahead_slots=num_lookahead_slots, num_lookahead_slots=num_lookahead_slots,
num_computed_slots=seq.data.get_num_computed_tokens(), num_computed_slots=seq.data.get_num_computed_tokens(),
) )
# Return any new copy-on-writes. # Return any new copy-on-writes.
new_cows = self.block_allocator.clear_copy_on_writes() new_cows = self.block_allocator.clear_copy_on_writes()
return new_cows return new_cows
def free(self, seq: Sequence) -> None: def free(self, seq: Sequence) -> None:
seq_id = seq.seq_id seq_id = seq.seq_id
if seq_id not in self.block_tables: if seq_id not in self.block_tables:
# Already freed or haven't been scheduled yet. # Already freed or haven't been scheduled yet.
return return
# Update seq block ids with the latest access time # Update seq block ids with the latest access time
self._last_access_blocks_tracker.update_seq_blocks_last_access( self._last_access_blocks_tracker.update_seq_blocks_last_access(
seq_id, self.block_tables[seq.seq_id].physical_block_ids) seq_id, self.block_tables[seq.seq_id].physical_block_ids)
# Untrack seq # Untrack seq
self._last_access_blocks_tracker.remove_seq(seq_id) self._last_access_blocks_tracker.remove_seq(seq_id)
self._computed_blocks_tracker.remove_seq(seq_id) self._computed_blocks_tracker.remove_seq(seq_id)
# Free table/blocks # Free table/blocks
self.block_tables[seq_id].free() self.block_tables[seq_id].free()
del self.block_tables[seq_id] del self.block_tables[seq_id]
def free_cross(self, seq_group: SequenceGroup) -> None: def free_cross(self, seq_group: SequenceGroup) -> None:
request_id = seq_group.request_id request_id = seq_group.request_id
if request_id not in self.cross_block_tables: if request_id not in self.cross_block_tables:
# Already freed or hasn't been scheduled yet. # Already freed or hasn't been scheduled yet.
return return
self.cross_block_tables[request_id].free() self.cross_block_tables[request_id].free()
del self.cross_block_tables[request_id] del self.cross_block_tables[request_id]
def get_block_table(self, seq: Sequence) -> List[int]: def get_block_table(self, seq: Sequence) -> List[int]:
block_ids = self.block_tables[seq.seq_id].physical_block_ids block_ids = self.block_tables[seq.seq_id].physical_block_ids
return block_ids # type: ignore return block_ids # type: ignore
def get_cross_block_table(self, seq_group: SequenceGroup) -> List[int]: def get_cross_block_table(self, seq_group: SequenceGroup) -> List[int]:
request_id = seq_group.request_id request_id = seq_group.request_id
assert request_id in self.cross_block_tables assert request_id in self.cross_block_tables
block_ids = self.cross_block_tables[request_id].physical_block_ids block_ids = self.cross_block_tables[request_id].physical_block_ids
assert all(b is not None for b in block_ids) assert all(b is not None for b in block_ids)
return block_ids # type: ignore return block_ids # type: ignore
def access_all_blocks_in_seq(self, seq: Sequence, now: float): def access_all_blocks_in_seq(self, seq: Sequence, now: float):
if self.enable_caching: if self.enable_caching:
# Record the latest access time for the sequence. The actual update # Record the latest access time for the sequence. The actual update
# of the block ids is deferred to the sequence free(..) call, since # of the block ids is deferred to the sequence free(..) call, since
# only during freeing of block ids, the blocks are actually added to # only during freeing of block ids, the blocks are actually added to
# the evictor (which is when the most updated time is required) # the evictor (which is when the most updated time is required)
# (This avoids expensive calls to mark_blocks_as_accessed(..)) # (This avoids expensive calls to mark_blocks_as_accessed(..))
self._last_access_blocks_tracker.update_last_access( self._last_access_blocks_tracker.update_last_access(
seq.seq_id, now) seq.seq_id, now)
def mark_blocks_as_computed(self, seq_group: SequenceGroup, def mark_blocks_as_computed(self, seq_group: SequenceGroup,
token_chunk_size: int): token_chunk_size: int):
# If prefix caching is enabled, mark immutable blocks as computed # If prefix caching is enabled, mark immutable blocks as computed
# right after they have been scheduled (for prefill). This assumes # right after they have been scheduled (for prefill). This assumes
# the scheduler is synchronous so blocks are actually computed when # the scheduler is synchronous so blocks are actually computed when
# scheduling the next batch. # scheduling the next batch.
self.block_allocator.mark_blocks_as_computed([]) self.block_allocator.mark_blocks_as_computed([])
def get_common_computed_block_ids( def get_common_computed_block_ids(
self, seqs: List[Sequence]) -> GenericSequence[int]: self, seqs: List[Sequence]) -> GenericSequence[int]:
"""Determine which blocks for which we skip prefill. """Determine which blocks for which we skip prefill.
With prefix caching we can skip prefill for previously-generated blocks. With prefix caching we can skip prefill for previously-generated blocks.
Currently, the attention implementation only supports skipping cached Currently, the attention implementation only supports skipping cached
blocks if they are a contiguous prefix of cached blocks. blocks if they are a contiguous prefix of cached blocks.
This method determines which blocks can be safely skipped for all This method determines which blocks can be safely skipped for all
sequences in the sequence group. sequences in the sequence group.
""" """
computed_seq_block_ids = [] computed_seq_block_ids = []
for seq in seqs: for seq in seqs:
computed_seq_block_ids.append( computed_seq_block_ids.append(
self._computed_blocks_tracker. self._computed_blocks_tracker.
get_cached_computed_blocks_and_update( get_cached_computed_blocks_and_update(
seq.seq_id, seq.seq_id,
self.block_tables[seq.seq_id].physical_block_ids)) self.block_tables[seq.seq_id].physical_block_ids))
# NOTE(sang): This assumes seq_block_ids doesn't contain any None. # NOTE(sang): This assumes seq_block_ids doesn't contain any None.
return self.block_allocator.get_common_computed_block_ids( return self.block_allocator.get_common_computed_block_ids(
computed_seq_block_ids) # type: ignore computed_seq_block_ids) # type: ignore
@@ -583,187 +583,187 @@ class BlockSpaceManagerV2(BlockSpaceManager):
raise TypeError(f"Unsupported multimodal namespace value type {type(value)}") raise TypeError(f"Unsupported multimodal namespace value type {type(value)}")
def fork(self, parent_seq: Sequence, child_seq: Sequence) -> None: def fork(self, parent_seq: Sequence, child_seq: Sequence) -> None:
if parent_seq.seq_id not in self.block_tables: if parent_seq.seq_id not in self.block_tables:
# Parent sequence has either been freed or never existed. # Parent sequence has either been freed or never existed.
return return
src_block_table = self.block_tables[parent_seq.seq_id] src_block_table = self.block_tables[parent_seq.seq_id]
self.block_tables[child_seq.seq_id] = src_block_table.fork() self.block_tables[child_seq.seq_id] = src_block_table.fork()
# Track child seq # Track child seq
self._computed_blocks_tracker.add_seq(child_seq.seq_id) self._computed_blocks_tracker.add_seq(child_seq.seq_id)
self._last_access_blocks_tracker.add_seq(child_seq.seq_id) self._last_access_blocks_tracker.add_seq(child_seq.seq_id)
def can_swap_in(self, seq_group: SequenceGroup, def can_swap_in(self, seq_group: SequenceGroup,
num_lookahead_slots: int) -> AllocStatus: num_lookahead_slots: int) -> AllocStatus:
"""Returns the AllocStatus for the given sequence_group """Returns the AllocStatus for the given sequence_group
with num_lookahead_slots. with num_lookahead_slots.
Args: Args:
sequence_group (SequenceGroup): The sequence group to swap in. sequence_group (SequenceGroup): The sequence group to swap in.
num_lookahead_slots (int): Number of lookahead slots used in num_lookahead_slots (int): Number of lookahead slots used in
speculative decoding, default to 0. speculative decoding, default to 0.
Returns: Returns:
AllocStatus: The AllocStatus for the given sequence group. AllocStatus: The AllocStatus for the given sequence group.
""" """
if self.block_allocator.content_offload_enabled: if self.block_allocator.content_offload_enabled:
return AllocStatus.NEVER return AllocStatus.NEVER
return self._can_swap(seq_group, Device.GPU, SequenceStatus.SWAPPED, return self._can_swap(seq_group, Device.GPU, SequenceStatus.SWAPPED,
num_lookahead_slots) num_lookahead_slots)
def swap_in(self, seq_group: SequenceGroup) -> List[Tuple[int, int]]: def swap_in(self, seq_group: SequenceGroup) -> List[Tuple[int, int]]:
"""Returns the block id mapping (from CPU to GPU) generated by """Returns the block id mapping (from CPU to GPU) generated by
swapping in the given seq_group with num_lookahead_slots. swapping in the given seq_group with num_lookahead_slots.
Args: Args:
seq_group (SequenceGroup): The sequence group to swap in. seq_group (SequenceGroup): The sequence group to swap in.
Returns: Returns:
List[Tuple[int, int]]: The mapping of swapping block from CPU List[Tuple[int, int]]: The mapping of swapping block from CPU
to GPU. to GPU.
""" """
physical_block_id_mapping = [] physical_block_id_mapping = []
for seq in seq_group.get_seqs(status=SequenceStatus.SWAPPED): for seq in seq_group.get_seqs(status=SequenceStatus.SWAPPED):
blocks = self.block_tables[seq.seq_id].blocks blocks = self.block_tables[seq.seq_id].blocks
if len(blocks) == 0: if len(blocks) == 0:
continue continue
seq_swap_mapping = self.block_allocator.swap(blocks=blocks, seq_swap_mapping = self.block_allocator.swap(blocks=blocks,
src_device=Device.CPU, src_device=Device.CPU,
dst_device=Device.GPU) dst_device=Device.GPU)
# Refresh the block ids of the table (post-swap) # Refresh the block ids of the table (post-swap)
self.block_tables[seq.seq_id].update(blocks) self.block_tables[seq.seq_id].update(blocks)
seq_physical_block_id_mapping = { seq_physical_block_id_mapping = {
self.block_allocator.get_physical_block_id( self.block_allocator.get_physical_block_id(
Device.CPU, cpu_block_id): Device.CPU, cpu_block_id):
self.block_allocator.get_physical_block_id( self.block_allocator.get_physical_block_id(
Device.GPU, gpu_block_id) Device.GPU, gpu_block_id)
for cpu_block_id, gpu_block_id in seq_swap_mapping.items() for cpu_block_id, gpu_block_id in seq_swap_mapping.items()
} }
physical_block_id_mapping.extend( physical_block_id_mapping.extend(
list(seq_physical_block_id_mapping.items())) list(seq_physical_block_id_mapping.items()))
return physical_block_id_mapping return physical_block_id_mapping
def can_swap_out(self, seq_group: SequenceGroup) -> bool: def can_swap_out(self, seq_group: SequenceGroup) -> bool:
"""Returns whether we can swap out the given sequence_group """Returns whether we can swap out the given sequence_group
with num_lookahead_slots. with num_lookahead_slots.
Args: Args:
seq_group (SequenceGroup): The sequence group to swap in. seq_group (SequenceGroup): The sequence group to swap in.
num_lookahead_slots (int): Number of lookahead slots used in num_lookahead_slots (int): Number of lookahead slots used in
speculative decoding, default to 0. speculative decoding, default to 0.
Returns: Returns:
bool: Whether it's possible to swap out current sequence group. bool: Whether it's possible to swap out current sequence group.
""" """
if self.block_allocator.content_offload_enabled: if self.block_allocator.content_offload_enabled:
return False return False
alloc_status = self._can_swap(seq_group, Device.CPU, alloc_status = self._can_swap(seq_group, Device.CPU,
SequenceStatus.RUNNING) SequenceStatus.RUNNING)
return alloc_status == AllocStatus.OK return alloc_status == AllocStatus.OK
def swap_out(self, seq_group: SequenceGroup) -> List[Tuple[int, int]]: def swap_out(self, seq_group: SequenceGroup) -> List[Tuple[int, int]]:
"""Returns the block id mapping (from GPU to CPU) generated by """Returns the block id mapping (from GPU to CPU) generated by
swapping out the given sequence_group with num_lookahead_slots. swapping out the given sequence_group with num_lookahead_slots.
Args: Args:
sequence_group (SequenceGroup): The sequence group to swap in. sequence_group (SequenceGroup): The sequence group to swap in.
Returns: Returns:
List[Tuple[int, int]]: The mapping of swapping block from List[Tuple[int, int]]: The mapping of swapping block from
GPU to CPU. GPU to CPU.
""" """
physical_block_id_mapping = [] physical_block_id_mapping = []
for seq in seq_group.get_seqs(status=SequenceStatus.RUNNING): for seq in seq_group.get_seqs(status=SequenceStatus.RUNNING):
blocks = self.block_tables[seq.seq_id].blocks blocks = self.block_tables[seq.seq_id].blocks
if len(blocks) == 0: if len(blocks) == 0:
continue continue
seq_swap_mapping = self.block_allocator.swap(blocks=blocks, seq_swap_mapping = self.block_allocator.swap(blocks=blocks,
src_device=Device.GPU, src_device=Device.GPU,
dst_device=Device.CPU) dst_device=Device.CPU)
# Refresh the block ids of the table (post-swap) # Refresh the block ids of the table (post-swap)
self.block_tables[seq.seq_id].update(blocks) self.block_tables[seq.seq_id].update(blocks)
seq_physical_block_id_mapping = { seq_physical_block_id_mapping = {
self.block_allocator.get_physical_block_id( self.block_allocator.get_physical_block_id(
Device.GPU, gpu_block_id): Device.GPU, gpu_block_id):
self.block_allocator.get_physical_block_id( self.block_allocator.get_physical_block_id(
Device.CPU, cpu_block_id) Device.CPU, cpu_block_id)
for gpu_block_id, cpu_block_id in seq_swap_mapping.items() for gpu_block_id, cpu_block_id in seq_swap_mapping.items()
} }
physical_block_id_mapping.extend( physical_block_id_mapping.extend(
list(seq_physical_block_id_mapping.items())) list(seq_physical_block_id_mapping.items()))
return physical_block_id_mapping return physical_block_id_mapping
def get_num_free_gpu_blocks(self) -> int: def get_num_free_gpu_blocks(self) -> int:
return self.block_allocator.get_num_free_blocks(Device.GPU) return self.block_allocator.get_num_free_blocks(Device.GPU)
def get_num_free_cpu_blocks(self) -> int: def get_num_free_cpu_blocks(self) -> int:
return self.block_allocator.get_num_free_blocks(Device.CPU) return self.block_allocator.get_num_free_blocks(Device.CPU)
def get_prefix_cache_hit_rate(self, device: Device) -> float: def get_prefix_cache_hit_rate(self, device: Device) -> float:
return self.block_allocator.get_prefix_cache_hit_rate(device) return self.block_allocator.get_prefix_cache_hit_rate(device)
def _can_swap(self, def _can_swap(self,
seq_group: SequenceGroup, seq_group: SequenceGroup,
device: Device, device: Device,
status: SequenceStatus, status: SequenceStatus,
num_lookahead_slots: int = 0) -> AllocStatus: num_lookahead_slots: int = 0) -> AllocStatus:
"""Returns the AllocStatus for swapping in/out the given sequence_group """Returns the AllocStatus for swapping in/out the given sequence_group
on to the 'device'. on to the 'device'.
Args: Args:
sequence_group (SequenceGroup): The sequence group to swap in. sequence_group (SequenceGroup): The sequence group to swap in.
device (Device): device to swap the 'seq_group' on. device (Device): device to swap the 'seq_group' on.
status (SequenceStatus): The status of sequence which is needed status (SequenceStatus): The status of sequence which is needed
for action. RUNNING for swap out and SWAPPED for swap in for action. RUNNING for swap out and SWAPPED for swap in
num_lookahead_slots (int): Number of lookahead slots used in num_lookahead_slots (int): Number of lookahead slots used in
speculative decoding, default to 0. speculative decoding, default to 0.
Returns: Returns:
AllocStatus: The AllocStatus for swapping in/out the given AllocStatus: The AllocStatus for swapping in/out the given
sequence_group on to the 'device'. sequence_group on to the 'device'.
""" """
# First determine the number of blocks that will be touched by this # First determine the number of blocks that will be touched by this
# swap. Then verify if there are available blocks in the device # swap. Then verify if there are available blocks in the device
# to perform the swap. # to perform the swap.
num_blocks_touched = 0 num_blocks_touched = 0
blocks: List[Block] = [] blocks: List[Block] = []
for seq in seq_group.get_seqs(status=status): for seq in seq_group.get_seqs(status=status):
block_table = self.block_tables[seq.seq_id] block_table = self.block_tables[seq.seq_id]
if block_table.blocks is not None: if block_table.blocks is not None:
# Compute the number blocks to touch for the tokens to be # Compute the number blocks to touch for the tokens to be
# appended. This does NOT include the full blocks that need # appended. This does NOT include the full blocks that need
# to be touched for the swap. # to be touched for the swap.
num_blocks_touched += \ num_blocks_touched += \
block_table.get_num_blocks_touched_by_append_slots( block_table.get_num_blocks_touched_by_append_slots(
block_table.get_unseen_token_ids(seq.get_token_ids()), block_table.get_unseen_token_ids(seq.get_token_ids()),
num_lookahead_slots=num_lookahead_slots) num_lookahead_slots=num_lookahead_slots)
blocks.extend(block_table.blocks) blocks.extend(block_table.blocks)
# Compute the number of full blocks to touch and add it to the # Compute the number of full blocks to touch and add it to the
# existing count of blocks to touch. # existing count of blocks to touch.
num_blocks_touched += self.block_allocator.get_num_full_blocks_touched( num_blocks_touched += self.block_allocator.get_num_full_blocks_touched(
blocks, device=device) blocks, device=device)
watermark_blocks = 0 watermark_blocks = 0
if device == Device.GPU: if device == Device.GPU:
watermark_blocks = self.watermark_blocks watermark_blocks = self.watermark_blocks
if self.block_allocator.get_num_total_blocks( if self.block_allocator.get_num_total_blocks(
device) < num_blocks_touched: device) < num_blocks_touched:
return AllocStatus.NEVER return AllocStatus.NEVER
elif self.block_allocator.get_num_free_blocks( elif self.block_allocator.get_num_free_blocks(
device) - num_blocks_touched >= watermark_blocks: device) - num_blocks_touched >= watermark_blocks:
return AllocStatus.OK return AllocStatus.OK
else: else:
return AllocStatus.LATER return AllocStatus.LATER

View File

@@ -7,128 +7,128 @@ from typing import Dict, List, OrderedDict, Tuple
ContentHash = bytes ContentHash = bytes
class EvictionPolicy(enum.Enum): class EvictionPolicy(enum.Enum):
"""Enum for eviction policy used by make_evictor to instantiate the correct """Enum for eviction policy used by make_evictor to instantiate the correct
Evictor subclass. Evictor subclass.
""" """
LRU = enum.auto() LRU = enum.auto()
FREQUENCY_AWARE = enum.auto() FREQUENCY_AWARE = enum.auto()
class Evictor(ABC): class Evictor(ABC):
"""The Evictor subclasses should be used by the BlockAllocator class to """The Evictor subclasses should be used by the BlockAllocator class to
handle eviction of freed PhysicalTokenBlocks. handle eviction of freed PhysicalTokenBlocks.
""" """
@abstractmethod @abstractmethod
def __init__(self): def __init__(self):
pass pass
@abstractmethod @abstractmethod
def __contains__(self, block_id: int) -> bool: def __contains__(self, block_id: int) -> bool:
pass pass
@abstractmethod @abstractmethod
def evict(self) -> Tuple[int, ContentHash]: def evict(self) -> Tuple[int, ContentHash]:
"""Runs the eviction algorithm and returns the evicted block's """Runs the eviction algorithm and returns the evicted block's
content hash along with physical block id along with physical block id content hash along with physical block id along with physical block id
""" """
pass pass
@abstractmethod @abstractmethod
def add(self, block_id: int, content_hash: ContentHash, def add(self, block_id: int, content_hash: ContentHash,
num_hashed_tokens: int, num_hashed_tokens: int,
last_accessed: float): last_accessed: float):
"""Adds block to the evictor, making it a candidate for eviction""" """Adds block to the evictor, making it a candidate for eviction"""
pass pass
@abstractmethod @abstractmethod
def update(self, block_id: int, last_accessed: float): def update(self, block_id: int, last_accessed: float):
"""Update corresponding block's access time in metadata""" """Update corresponding block's access time in metadata"""
pass pass
@abstractmethod @abstractmethod
def remove(self, block_id: int): def remove(self, block_id: int):
"""Remove a given block id from the cache.""" """Remove a given block id from the cache."""
pass pass
@property @property
@abstractmethod @abstractmethod
def num_blocks(self) -> int: def num_blocks(self) -> int:
pass pass
class BlockMetaData(): class BlockMetaData():
"""Data structure for storing key data describe cached block, so that """Data structure for storing key data describe cached block, so that
evitor could use to make its decision which one to choose for eviction evitor could use to make its decision which one to choose for eviction
Here we use physical block id as the dict key, as there maybe several Here we use physical block id as the dict key, as there maybe several
blocks with the same content hash, but their physical id is unique. blocks with the same content hash, but their physical id is unique.
""" """
def __init__(self, content_hash: ContentHash, num_hashed_tokens: int, def __init__(self, content_hash: ContentHash, num_hashed_tokens: int,
last_accessed: float): last_accessed: float):
self.content_hash = content_hash self.content_hash = content_hash
self.num_hashed_tokens = num_hashed_tokens self.num_hashed_tokens = num_hashed_tokens
self.last_accessed = last_accessed self.last_accessed = last_accessed
class LRUEvictor(Evictor): class LRUEvictor(Evictor):
"""Evicts in a least-recently-used order using the last_accessed timestamp """Evicts in a least-recently-used order using the last_accessed timestamp
that's recorded in the PhysicalTokenBlock. If there are multiple blocks with that's recorded in the PhysicalTokenBlock. If there are multiple blocks with
the same last_accessed time, then the one with the largest num_hashed_tokens the same last_accessed time, then the one with the largest num_hashed_tokens
will be evicted. If two blocks each have the lowest last_accessed time and will be evicted. If two blocks each have the lowest last_accessed time and
highest num_hashed_tokens value, then one will be chose arbitrarily highest num_hashed_tokens value, then one will be chose arbitrarily
""" """
def __init__(self): def __init__(self):
self.free_table: OrderedDict[int, BlockMetaData] = OrderedDict() self.free_table: OrderedDict[int, BlockMetaData] = OrderedDict()
def __contains__(self, block_id: int) -> bool: def __contains__(self, block_id: int) -> bool:
return block_id in self.free_table return block_id in self.free_table
def evict(self) -> Tuple[int, ContentHash]: def evict(self) -> Tuple[int, ContentHash]:
if len(self.free_table) == 0: if len(self.free_table) == 0:
raise ValueError("No usable cache memory left") raise ValueError("No usable cache memory left")
evicted_block, evicted_block_id = None, None evicted_block, evicted_block_id = None, None
# The blocks with the lowest timestamps should be placed consecutively # The blocks with the lowest timestamps should be placed consecutively
# at the start of OrderedDict. Loop through all these blocks to # at the start of OrderedDict. Loop through all these blocks to
# find the one with maximum number of hashed tokens. # find the one with maximum number of hashed tokens.
for _id, block in self.free_table.items(): for _id, block in self.free_table.items():
if evicted_block is None: if evicted_block is None:
evicted_block, evicted_block_id = block, _id evicted_block, evicted_block_id = block, _id
continue continue
if evicted_block.last_accessed < block.last_accessed: if evicted_block.last_accessed < block.last_accessed:
break break
if evicted_block.num_hashed_tokens < block.num_hashed_tokens: if evicted_block.num_hashed_tokens < block.num_hashed_tokens:
evicted_block, evicted_block_id = block, _id evicted_block, evicted_block_id = block, _id
assert evicted_block is not None assert evicted_block is not None
assert evicted_block_id is not None assert evicted_block_id is not None
self.free_table.pop(evicted_block_id) self.free_table.pop(evicted_block_id)
return evicted_block_id, evicted_block.content_hash return evicted_block_id, evicted_block.content_hash
def add(self, block_id: int, content_hash: ContentHash, def add(self, block_id: int, content_hash: ContentHash,
num_hashed_tokens: int, num_hashed_tokens: int,
last_accessed: float): last_accessed: float):
self.free_table[block_id] = BlockMetaData(content_hash, self.free_table[block_id] = BlockMetaData(content_hash,
num_hashed_tokens, num_hashed_tokens,
last_accessed) last_accessed)
def update(self, block_id: int, last_accessed: float): def update(self, block_id: int, last_accessed: float):
self.free_table[block_id].last_accessed = last_accessed self.free_table[block_id].last_accessed = last_accessed
def remove(self, block_id: int): def remove(self, block_id: int):
if block_id not in self.free_table: if block_id not in self.free_table:
raise ValueError( raise ValueError(
"Attempting to remove block that's not in the evictor") "Attempting to remove block that's not in the evictor")
self.free_table.pop(block_id) self.free_table.pop(block_id)
@property @property
def num_blocks(self) -> int: def num_blocks(self) -> int:
return len(self.free_table) return len(self.free_table)
@@ -261,8 +261,8 @@ def eviction_policy_from_env(
raise ValueError( raise ValueError(
"BI100_KV_EVICTION_POLICY must be one of: frequency, lru") "BI100_KV_EVICTION_POLICY must be one of: frequency, lru")
return policies[value] return policies[value]
def make_evictor(eviction_policy: EvictionPolicy) -> Evictor: def make_evictor(eviction_policy: EvictionPolicy) -> Evictor:
if eviction_policy == EvictionPolicy.LRU: if eviction_policy == EvictionPolicy.LRU:
return LRUEvictor() return LRUEvictor()

View File

@@ -1,391 +1,391 @@
"""Sampling parameters for text generation.""" """Sampling parameters for text generation."""
import copy import copy
from dataclasses import dataclass from dataclasses import dataclass
from enum import Enum, IntEnum from enum import Enum, IntEnum
from functools import cached_property from functools import cached_property
from typing import Any, Callable, Dict, List, Optional, Set, Union from typing import Any, Callable, Dict, List, Optional, Set, Union
import msgspec import msgspec
import torch import torch
from pydantic import BaseModel from pydantic import BaseModel
from typing_extensions import Annotated from typing_extensions import Annotated
from vllm.logger import init_logger from vllm.logger import init_logger
logger = init_logger(__name__) logger = init_logger(__name__)
_SAMPLING_EPS = 1e-5 _SAMPLING_EPS = 1e-5
_MAX_TEMP = 1e-2 _MAX_TEMP = 1e-2
class SamplingType(IntEnum): class SamplingType(IntEnum):
GREEDY = 0 GREEDY = 0
RANDOM = 1 RANDOM = 1
RANDOM_SEED = 2 RANDOM_SEED = 2
LogitsProcessor = Union[Callable[[List[int], torch.Tensor], torch.Tensor], LogitsProcessor = Union[Callable[[List[int], torch.Tensor], torch.Tensor],
Callable[[List[int], List[int], torch.Tensor], Callable[[List[int], List[int], torch.Tensor],
torch.Tensor]] torch.Tensor]]
"""LogitsProcessor is a function that takes a list """LogitsProcessor is a function that takes a list
of previously generated tokens, the logits tensor of previously generated tokens, the logits tensor
for the next token and, optionally, prompt tokens as a for the next token and, optionally, prompt tokens as a
first argument, and returns a modified tensor of logits first argument, and returns a modified tensor of logits
to sample from.""" to sample from."""
# maybe make msgspec? # maybe make msgspec?
@dataclass @dataclass
class GuidedDecodingParams: class GuidedDecodingParams:
"""One of these fields will be used to build a logit processor.""" """One of these fields will be used to build a logit processor."""
json: Optional[Union[str, Dict]] = None json: Optional[Union[str, Dict]] = None
regex: Optional[str] = None regex: Optional[str] = None
choice: Optional[List[str]] = None choice: Optional[List[str]] = None
grammar: Optional[str] = None grammar: Optional[str] = None
json_object: Optional[bool] = None json_object: Optional[bool] = None
"""These are other options that can be set""" """These are other options that can be set"""
backend: Optional[str] = None backend: Optional[str] = None
whitespace_pattern: Optional[str] = None whitespace_pattern: Optional[str] = None
@staticmethod @staticmethod
def from_optional( def from_optional(
json: Optional[Union[Dict, BaseModel, str]], json: Optional[Union[Dict, BaseModel, str]],
regex: Optional[str] = None, regex: Optional[str] = None,
choice: Optional[List[str]] = None, choice: Optional[List[str]] = None,
grammar: Optional[str] = None, grammar: Optional[str] = None,
json_object: Optional[bool] = None, json_object: Optional[bool] = None,
backend: Optional[str] = None, backend: Optional[str] = None,
whitespace_pattern: Optional[str] = None, whitespace_pattern: Optional[str] = None,
) -> "GuidedDecodingParams": ) -> "GuidedDecodingParams":
# Extract json schemas from pydantic models # Extract json schemas from pydantic models
if isinstance(json, (BaseModel, type(BaseModel))): if isinstance(json, (BaseModel, type(BaseModel))):
json = json.model_json_schema() json = json.model_json_schema()
return GuidedDecodingParams( return GuidedDecodingParams(
json=json, json=json,
regex=regex, regex=regex,
choice=choice, choice=choice,
grammar=grammar, grammar=grammar,
json_object=json_object, json_object=json_object,
backend=backend, backend=backend,
whitespace_pattern=whitespace_pattern, whitespace_pattern=whitespace_pattern,
) )
def __post_init__(self): def __post_init__(self):
"""Validate that some fields are mutually exclusive.""" """Validate that some fields are mutually exclusive."""
guide_count = sum([ guide_count = sum([
self.json is not None, self.regex is not None, self.choice self.json is not None, self.regex is not None, self.choice
is not None, self.grammar is not None, self.json_object is not None is not None, self.grammar is not None, self.json_object is not None
]) ])
if guide_count > 1: if guide_count > 1:
raise ValueError( raise ValueError(
"You can only use one kind of guided decoding but multiple are " "You can only use one kind of guided decoding but multiple are "
f"specified: {self.__dict__}") f"specified: {self.__dict__}")
class RequestOutputKind(Enum): class RequestOutputKind(Enum):
# Return entire output so far in every RequestOutput # Return entire output so far in every RequestOutput
CUMULATIVE = 0 CUMULATIVE = 0
# Return only deltas in each RequestOutput # Return only deltas in each RequestOutput
DELTA = 1 DELTA = 1
# Do not return intermediate RequestOuputs # Do not return intermediate RequestOuputs
FINAL_ONLY = 2 FINAL_ONLY = 2
class SamplingParams( class SamplingParams(
msgspec.Struct, msgspec.Struct,
omit_defaults=True, # type: ignore[call-arg] omit_defaults=True, # type: ignore[call-arg]
# required for @cached_property. # required for @cached_property.
dict=True): # type: ignore[call-arg] dict=True): # type: ignore[call-arg]
"""Sampling parameters for text generation. """Sampling parameters for text generation.
Overall, we follow the sampling parameters from the OpenAI text completion Overall, we follow the sampling parameters from the OpenAI text completion
API (https://platform.openai.com/docs/api-reference/completions/create). API (https://platform.openai.com/docs/api-reference/completions/create).
In addition, we support beam search, which is not supported by OpenAI. In addition, we support beam search, which is not supported by OpenAI.
Args: Args:
n: Number of output sequences to return for the given prompt. n: Number of output sequences to return for the given prompt.
best_of: Number of output sequences that are generated from the prompt. best_of: Number of output sequences that are generated from the prompt.
From these `best_of` sequences, the top `n` sequences are returned. From these `best_of` sequences, the top `n` sequences are returned.
`best_of` must be greater than or equal to `n`. By default, `best_of` must be greater than or equal to `n`. By default,
`best_of` is set to `n`. `best_of` is set to `n`.
presence_penalty: Float that penalizes new tokens based on whether they presence_penalty: Float that penalizes new tokens based on whether they
appear in the generated text so far. Values > 0 encourage the model appear in the generated text so far. Values > 0 encourage the model
to use new tokens, while values < 0 encourage the model to repeat to use new tokens, while values < 0 encourage the model to repeat
tokens. tokens.
frequency_penalty: Float that penalizes new tokens based on their frequency_penalty: Float that penalizes new tokens based on their
frequency in the generated text so far. Values > 0 encourage the frequency in the generated text so far. Values > 0 encourage the
model to use new tokens, while values < 0 encourage the model to model to use new tokens, while values < 0 encourage the model to
repeat tokens. repeat tokens.
repetition_penalty: Float that penalizes new tokens based on whether repetition_penalty: Float that penalizes new tokens based on whether
they appear in the prompt and the generated text so far. Values > 1 they appear in the prompt and the generated text so far. Values > 1
encourage the model to use new tokens, while values < 1 encourage encourage the model to use new tokens, while values < 1 encourage
the model to repeat tokens. the model to repeat tokens.
temperature: Float that controls the randomness of the sampling. Lower temperature: Float that controls the randomness of the sampling. Lower
values make the model more deterministic, while higher values make values make the model more deterministic, while higher values make
the model more random. Zero means greedy sampling. the model more random. Zero means greedy sampling.
top_p: Float that controls the cumulative probability of the top tokens top_p: Float that controls the cumulative probability of the top tokens
to consider. Must be in (0, 1]. Set to 1 to consider all tokens. to consider. Must be in (0, 1]. Set to 1 to consider all tokens.
top_k: Integer that controls the number of top tokens to consider. Set top_k: Integer that controls the number of top tokens to consider. Set
to -1 to consider all tokens. to -1 to consider all tokens.
min_p: Float that represents the minimum probability for a token to be min_p: Float that represents the minimum probability for a token to be
considered, relative to the probability of the most likely token. considered, relative to the probability of the most likely token.
Must be in [0, 1]. Set to 0 to disable this. Must be in [0, 1]. Set to 0 to disable this.
seed: Random seed to use for the generation. seed: Random seed to use for the generation.
stop: List of strings that stop the generation when they are generated. stop: List of strings that stop the generation when they are generated.
The returned output will not contain the stop strings. The returned output will not contain the stop strings.
stop_token_ids: List of tokens that stop the generation when they are stop_token_ids: List of tokens that stop the generation when they are
generated. The returned output will contain the stop tokens unless generated. The returned output will contain the stop tokens unless
the stop tokens are special tokens. the stop tokens are special tokens.
include_stop_str_in_output: Whether to include the stop strings in include_stop_str_in_output: Whether to include the stop strings in
output text. Defaults to False. output text. Defaults to False.
ignore_eos: Whether to ignore the EOS token and continue generating ignore_eos: Whether to ignore the EOS token and continue generating
tokens after the EOS token is generated. tokens after the EOS token is generated.
max_tokens: Maximum number of tokens to generate per output sequence. max_tokens: Maximum number of tokens to generate per output sequence.
min_tokens: Minimum number of tokens to generate per output sequence min_tokens: Minimum number of tokens to generate per output sequence
before EOS or stop_token_ids can be generated before EOS or stop_token_ids can be generated
logprobs: Number of log probabilities to return per output token. logprobs: Number of log probabilities to return per output token.
When set to None, no probability is returned. If set to a non-None When set to None, no probability is returned. If set to a non-None
value, the result includes the log probabilities of the specified value, the result includes the log probabilities of the specified
number of most likely tokens, as well as the chosen tokens. number of most likely tokens, as well as the chosen tokens.
Note that the implementation follows the OpenAI API: The API will Note that the implementation follows the OpenAI API: The API will
always return the log probability of the sampled token, so there always return the log probability of the sampled token, so there
may be up to `logprobs+1` elements in the response. may be up to `logprobs+1` elements in the response.
prompt_logprobs: Number of log probabilities to return per prompt token. prompt_logprobs: Number of log probabilities to return per prompt token.
detokenize: Whether to detokenize the output. Defaults to True. detokenize: Whether to detokenize the output. Defaults to True.
skip_special_tokens: Whether to skip special tokens in the output. skip_special_tokens: Whether to skip special tokens in the output.
spaces_between_special_tokens: Whether to add spaces between special spaces_between_special_tokens: Whether to add spaces between special
tokens in the output. Defaults to True. tokens in the output. Defaults to True.
logits_processors: List of functions that modify logits based on logits_processors: List of functions that modify logits based on
previously generated tokens, and optionally prompt tokens as previously generated tokens, and optionally prompt tokens as
a first argument. a first argument.
truncate_prompt_tokens: If set to an integer k, will use only the last k truncate_prompt_tokens: If set to an integer k, will use only the last k
tokens from the prompt (i.e., left truncation). Defaults to None tokens from the prompt (i.e., left truncation). Defaults to None
(i.e., no truncation). (i.e., no truncation).
guided_decoding: If provided, the engine will construct a guided guided_decoding: If provided, the engine will construct a guided
decoding logits processor from these parameters. Defaults to None. decoding logits processor from these parameters. Defaults to None.
logit_bias: If provided, the engine will construct a logits processor logit_bias: If provided, the engine will construct a logits processor
that applies these logit biases. Defaults to None. that applies these logit biases. Defaults to None.
allowed_token_ids: If provided, the engine will construct a logits allowed_token_ids: If provided, the engine will construct a logits
processor which only retains scores for the given token ids. processor which only retains scores for the given token ids.
Defaults to None. Defaults to None.
prompt_logprob_positions: Optional prompt-token positions whose logits prompt_logprob_positions: Optional prompt-token positions whose logits
should be materialized. None preserves the standard all-position should be materialized. None preserves the standard all-position
prompt-logprob behavior. prompt-logprob behavior.
""" """
n: int = 1 n: int = 1
best_of: Optional[int] = None best_of: Optional[int] = None
_real_n: Optional[int] = None _real_n: Optional[int] = None
presence_penalty: float = 0.0 presence_penalty: float = 0.0
frequency_penalty: float = 0.0 frequency_penalty: float = 0.0
repetition_penalty: float = 1.0 repetition_penalty: float = 1.0
temperature: float = 1.0 temperature: float = 1.0
top_p: float = 1.0 top_p: float = 1.0
top_k: int = -1 top_k: int = -1
min_p: float = 0.0 min_p: float = 0.0
seed: Optional[int] = None seed: Optional[int] = None
stop: Optional[Union[str, List[str]]] = None stop: Optional[Union[str, List[str]]] = None
stop_token_ids: Optional[List[int]] = None stop_token_ids: Optional[List[int]] = None
ignore_eos: bool = False ignore_eos: bool = False
max_tokens: Optional[int] = 16 max_tokens: Optional[int] = 16
min_tokens: int = 0 min_tokens: int = 0
logprobs: Optional[int] = None logprobs: Optional[int] = None
prompt_logprobs: Optional[int] = None prompt_logprobs: Optional[int] = None
# NOTE: This parameter is only exposed at the engine level for now. # NOTE: This parameter is only exposed at the engine level for now.
# It is not exposed in the OpenAI API server, as the OpenAI API does # It is not exposed in the OpenAI API server, as the OpenAI API does
# not support returning only a list of token IDs. # not support returning only a list of token IDs.
detokenize: bool = True detokenize: bool = True
skip_special_tokens: bool = True skip_special_tokens: bool = True
spaces_between_special_tokens: bool = True spaces_between_special_tokens: bool = True
# Optional[List[LogitsProcessor]] type. We use Any here because # Optional[List[LogitsProcessor]] type. We use Any here because
# Optional[List[LogitsProcessor]] type is not supported by msgspec. # Optional[List[LogitsProcessor]] type is not supported by msgspec.
logits_processors: Optional[Any] = None logits_processors: Optional[Any] = None
include_stop_str_in_output: bool = False include_stop_str_in_output: bool = False
truncate_prompt_tokens: Optional[Annotated[int, msgspec.Meta(ge=1)]] = None truncate_prompt_tokens: Optional[Annotated[int, msgspec.Meta(ge=1)]] = None
output_kind: RequestOutputKind = RequestOutputKind.CUMULATIVE output_kind: RequestOutputKind = RequestOutputKind.CUMULATIVE
# The below fields are not supposed to be used as an input. # The below fields are not supposed to be used as an input.
# They are set in post_init. # They are set in post_init.
output_text_buffer_length: int = 0 output_text_buffer_length: int = 0
_all_stop_token_ids: Set[int] = msgspec.field(default_factory=set) _all_stop_token_ids: Set[int] = msgspec.field(default_factory=set)
# Fields used to construct logits processors # Fields used to construct logits processors
guided_decoding: Optional[GuidedDecodingParams] = None guided_decoding: Optional[GuidedDecodingParams] = None
logit_bias: Optional[Dict[int, float]] = None logit_bias: Optional[Dict[int, float]] = None
allowed_token_ids: Optional[List[int]] = None allowed_token_ids: Optional[List[int]] = None
prompt_logprob_positions: Optional[List[int]] = None prompt_logprob_positions: Optional[List[int]] = None
@staticmethod @staticmethod
def from_optional( def from_optional(
n: Optional[int] = 1, n: Optional[int] = 1,
best_of: Optional[int] = None, best_of: Optional[int] = None,
presence_penalty: Optional[float] = 0.0, presence_penalty: Optional[float] = 0.0,
frequency_penalty: Optional[float] = 0.0, frequency_penalty: Optional[float] = 0.0,
repetition_penalty: Optional[float] = 1.0, repetition_penalty: Optional[float] = 1.0,
temperature: Optional[float] = 1.0, temperature: Optional[float] = 1.0,
top_p: Optional[float] = 1.0, top_p: Optional[float] = 1.0,
top_k: int = -1, top_k: int = -1,
min_p: float = 0.0, min_p: float = 0.0,
seed: Optional[int] = None, seed: Optional[int] = None,
stop: Optional[Union[str, List[str]]] = None, stop: Optional[Union[str, List[str]]] = None,
stop_token_ids: Optional[List[int]] = None, stop_token_ids: Optional[List[int]] = None,
include_stop_str_in_output: bool = False, include_stop_str_in_output: bool = False,
ignore_eos: bool = False, ignore_eos: bool = False,
max_tokens: Optional[int] = 16, max_tokens: Optional[int] = 16,
min_tokens: int = 0, min_tokens: int = 0,
logprobs: Optional[int] = None, logprobs: Optional[int] = None,
prompt_logprobs: Optional[int] = None, prompt_logprobs: Optional[int] = None,
detokenize: bool = True, detokenize: bool = True,
skip_special_tokens: bool = True, skip_special_tokens: bool = True,
spaces_between_special_tokens: bool = True, spaces_between_special_tokens: bool = True,
logits_processors: Optional[List[LogitsProcessor]] = None, logits_processors: Optional[List[LogitsProcessor]] = None,
truncate_prompt_tokens: Optional[Annotated[int, truncate_prompt_tokens: Optional[Annotated[int,
msgspec.Meta(ge=1)]] = None, msgspec.Meta(ge=1)]] = None,
output_kind: RequestOutputKind = RequestOutputKind.CUMULATIVE, output_kind: RequestOutputKind = RequestOutputKind.CUMULATIVE,
guided_decoding: Optional[GuidedDecodingParams] = None, guided_decoding: Optional[GuidedDecodingParams] = None,
logit_bias: Optional[Union[Dict[int, float], Dict[str, float]]] = None, logit_bias: Optional[Union[Dict[int, float], Dict[str, float]]] = None,
allowed_token_ids: Optional[List[int]] = None, allowed_token_ids: Optional[List[int]] = None,
prompt_logprob_positions: Optional[List[int]] = None, prompt_logprob_positions: Optional[List[int]] = None,
) -> "SamplingParams": ) -> "SamplingParams":
if logit_bias is not None: if logit_bias is not None:
logit_bias = { logit_bias = {
int(token): bias int(token): bias
for token, bias in logit_bias.items() for token, bias in logit_bias.items()
} }
return SamplingParams( return SamplingParams(
n=1 if n is None else n, n=1 if n is None else n,
best_of=best_of, best_of=best_of,
presence_penalty=0.0 presence_penalty=0.0
if presence_penalty is None else presence_penalty, if presence_penalty is None else presence_penalty,
frequency_penalty=0.0 frequency_penalty=0.0
if frequency_penalty is None else frequency_penalty, if frequency_penalty is None else frequency_penalty,
repetition_penalty=1.0 repetition_penalty=1.0
if repetition_penalty is None else repetition_penalty, if repetition_penalty is None else repetition_penalty,
temperature=1.0 if temperature is None else temperature, temperature=1.0 if temperature is None else temperature,
top_p=1.0 if top_p is None else top_p, top_p=1.0 if top_p is None else top_p,
top_k=top_k, top_k=top_k,
min_p=min_p, min_p=min_p,
seed=seed, seed=seed,
stop=stop, stop=stop,
stop_token_ids=stop_token_ids, stop_token_ids=stop_token_ids,
include_stop_str_in_output=include_stop_str_in_output, include_stop_str_in_output=include_stop_str_in_output,
ignore_eos=ignore_eos, ignore_eos=ignore_eos,
max_tokens=max_tokens, max_tokens=max_tokens,
min_tokens=min_tokens, min_tokens=min_tokens,
logprobs=logprobs, logprobs=logprobs,
prompt_logprobs=prompt_logprobs, prompt_logprobs=prompt_logprobs,
detokenize=detokenize, detokenize=detokenize,
skip_special_tokens=skip_special_tokens, skip_special_tokens=skip_special_tokens,
spaces_between_special_tokens=spaces_between_special_tokens, spaces_between_special_tokens=spaces_between_special_tokens,
logits_processors=logits_processors, logits_processors=logits_processors,
truncate_prompt_tokens=truncate_prompt_tokens, truncate_prompt_tokens=truncate_prompt_tokens,
output_kind=output_kind, output_kind=output_kind,
guided_decoding=guided_decoding, guided_decoding=guided_decoding,
logit_bias=logit_bias, logit_bias=logit_bias,
allowed_token_ids=allowed_token_ids, allowed_token_ids=allowed_token_ids,
prompt_logprob_positions=prompt_logprob_positions, prompt_logprob_positions=prompt_logprob_positions,
) )
def __post_init__(self) -> None: def __post_init__(self) -> None:
# how we deal with `best_of``: # how we deal with `best_of``:
# if `best_of`` is not set, we default to `n`; # if `best_of`` is not set, we default to `n`;
# if `best_of`` is set, we set `n`` to `best_of`, # if `best_of`` is set, we set `n`` to `best_of`,
# and set `_real_n`` to the original `n`. # and set `_real_n`` to the original `n`.
# when we return the result, we will check # when we return the result, we will check
# if we need to return `n` or `_real_n` results # if we need to return `n` or `_real_n` results
if self.best_of: if self.best_of:
if self.best_of < self.n: if self.best_of < self.n:
raise ValueError( raise ValueError(
f"best_of must be greater than or equal to n, " f"best_of must be greater than or equal to n, "
f"got n={self.n} and best_of={self.best_of}.") f"got n={self.n} and best_of={self.best_of}.")
self._real_n = self.n self._real_n = self.n
self.n = self.best_of self.n = self.best_of
if 0 < self.temperature < _MAX_TEMP: if 0 < self.temperature < _MAX_TEMP:
logger.warning( logger.warning(
"temperature %s is less than %s, which may cause numerical " "temperature %s is less than %s, which may cause numerical "
"errors nan or inf in tensors. We have maxed it out to %s.", "errors nan or inf in tensors. We have maxed it out to %s.",
self.temperature, _MAX_TEMP, _MAX_TEMP) self.temperature, _MAX_TEMP, _MAX_TEMP)
self.temperature = max(self.temperature, _MAX_TEMP) self.temperature = max(self.temperature, _MAX_TEMP)
if self.seed == -1: if self.seed == -1:
self.seed = None self.seed = None
else: else:
self.seed = self.seed self.seed = self.seed
if self.stop is None: if self.stop is None:
self.stop = [] self.stop = []
elif isinstance(self.stop, str): elif isinstance(self.stop, str):
self.stop = [self.stop] self.stop = [self.stop]
else: else:
self.stop = list(self.stop) self.stop = list(self.stop)
if self.stop_token_ids is None: if self.stop_token_ids is None:
self.stop_token_ids = [] self.stop_token_ids = []
else: else:
self.stop_token_ids = list(self.stop_token_ids) self.stop_token_ids = list(self.stop_token_ids)
self.logprobs = 1 if self.logprobs is True else self.logprobs self.logprobs = 1 if self.logprobs is True else self.logprobs
self.prompt_logprobs = (1 if self.prompt_logprobs is True else self.prompt_logprobs = (1 if self.prompt_logprobs is True else
self.prompt_logprobs) self.prompt_logprobs)
if self.prompt_logprob_positions is not None: if self.prompt_logprob_positions is not None:
self.prompt_logprob_positions = list( self.prompt_logprob_positions = list(
self.prompt_logprob_positions) self.prompt_logprob_positions)
# Number of characters to hold back for stop string evaluation # Number of characters to hold back for stop string evaluation
# until sequence is finished. # until sequence is finished.
if self.stop and not self.include_stop_str_in_output: if self.stop and not self.include_stop_str_in_output:
self.output_text_buffer_length = max(len(s) for s in self.stop) - 1 self.output_text_buffer_length = max(len(s) for s in self.stop) - 1
self._verify_args() self._verify_args()
if self.temperature < _SAMPLING_EPS: if self.temperature < _SAMPLING_EPS:
# Zero temperature means greedy sampling. # Zero temperature means greedy sampling.
self.top_p = 1.0 self.top_p = 1.0
self.top_k = -1 self.top_k = -1
self.min_p = 0.0 self.min_p = 0.0
self._verify_greedy_sampling() self._verify_greedy_sampling()
# eos_token_id is added to this by the engine # eos_token_id is added to this by the engine
self._all_stop_token_ids = set(self.stop_token_ids) self._all_stop_token_ids = set(self.stop_token_ids)
def _verify_args(self) -> None: def _verify_args(self) -> None:
if not isinstance(self.n, int): if not isinstance(self.n, int):
raise ValueError(f"n must be an int, but is of " raise ValueError(f"n must be an int, but is of "
f"type {type(self.n)}") f"type {type(self.n)}")
if self.n < 1: if self.n < 1:
raise ValueError(f"n must be at least 1, got {self.n}.") raise ValueError(f"n must be at least 1, got {self.n}.")
if not -2.0 <= self.presence_penalty <= 2.0: if not -2.0 <= self.presence_penalty <= 2.0:
raise ValueError("presence_penalty must be in [-2, 2], got " raise ValueError("presence_penalty must be in [-2, 2], got "
f"{self.presence_penalty}.") f"{self.presence_penalty}.")
if not -2.0 <= self.frequency_penalty <= 2.0: if not -2.0 <= self.frequency_penalty <= 2.0:
raise ValueError("frequency_penalty must be in [-2, 2], got " raise ValueError("frequency_penalty must be in [-2, 2], got "
f"{self.frequency_penalty}.") f"{self.frequency_penalty}.")
if not 0.0 < self.repetition_penalty <= 2.0: if not 0.0 < self.repetition_penalty <= 2.0:
raise ValueError("repetition_penalty must be in (0, 2], got " raise ValueError("repetition_penalty must be in (0, 2], got "
f"{self.repetition_penalty}.") f"{self.repetition_penalty}.")
if self.temperature < 0.0: if self.temperature < 0.0:
raise ValueError( raise ValueError(
f"temperature must be non-negative, got {self.temperature}.") f"temperature must be non-negative, got {self.temperature}.")
if not 0.0 < self.top_p <= 1.0: if not 0.0 < self.top_p <= 1.0:
raise ValueError(f"top_p must be in (0, 1], got {self.top_p}.") raise ValueError(f"top_p must be in (0, 1], got {self.top_p}.")
if self.top_k < -1 or self.top_k == 0: if self.top_k < -1 or self.top_k == 0:
raise ValueError(f"top_k must be -1 (disable), or at least 1, " raise ValueError(f"top_k must be -1 (disable), or at least 1, "
f"got {self.top_k}.") f"got {self.top_k}.")
if not isinstance(self.top_k, int): if not isinstance(self.top_k, int):
raise TypeError( raise TypeError(
f"top_k must be an integer, got {type(self.top_k).__name__}") f"top_k must be an integer, got {type(self.top_k).__name__}")
if not 0.0 <= self.min_p <= 1.0: if not 0.0 <= self.min_p <= 1.0:
raise ValueError("min_p must be in [0, 1], got " raise ValueError("min_p must be in [0, 1], got "
f"{self.min_p}.") f"{self.min_p}.")
if self.max_tokens is not None and self.max_tokens < 1: if self.max_tokens is not None and self.max_tokens < 1:
raise ValueError( raise ValueError(
f"max_tokens must be at least 1, got {self.max_tokens}.") f"max_tokens must be at least 1, got {self.max_tokens}.")
if self.min_tokens < 0: if self.min_tokens < 0:
raise ValueError(f"min_tokens must be greater than or equal to 0, " raise ValueError(f"min_tokens must be greater than or equal to 0, "
f"got {self.min_tokens}.") f"got {self.min_tokens}.")
if self.max_tokens is not None and self.min_tokens > self.max_tokens: if self.max_tokens is not None and self.min_tokens > self.max_tokens:
raise ValueError( raise ValueError(
f"min_tokens must be less than or equal to " f"min_tokens must be less than or equal to "
f"max_tokens={self.max_tokens}, got {self.min_tokens}.") f"max_tokens={self.max_tokens}, got {self.min_tokens}.")
if self.logprobs is not None and self.logprobs < 0: if self.logprobs is not None and self.logprobs < 0:
raise ValueError( raise ValueError(
f"logprobs must be non-negative, got {self.logprobs}.") f"logprobs must be non-negative, got {self.logprobs}.")
if self.prompt_logprobs is not None and self.prompt_logprobs < 0: if self.prompt_logprobs is not None and self.prompt_logprobs < 0:
raise ValueError(f"prompt_logprobs must be non-negative, got " raise ValueError(f"prompt_logprobs must be non-negative, got "
f"{self.prompt_logprobs}.") f"{self.prompt_logprobs}.")
@@ -407,114 +407,114 @@ class SamplingParams(
raise ValueError( raise ValueError(
"prompt_logprob_positions must be a sorted unique list " "prompt_logprob_positions must be a sorted unique list "
"of positive integers.") "of positive integers.")
if (self.truncate_prompt_tokens is not None if (self.truncate_prompt_tokens is not None
and self.truncate_prompt_tokens < 1): and self.truncate_prompt_tokens < 1):
raise ValueError(f"truncate_prompt_tokens must be >= 1, " raise ValueError(f"truncate_prompt_tokens must be >= 1, "
f"got {self.truncate_prompt_tokens}") f"got {self.truncate_prompt_tokens}")
assert isinstance(self.stop, list) assert isinstance(self.stop, list)
if any(not stop_str for stop_str in self.stop): if any(not stop_str for stop_str in self.stop):
raise ValueError("stop cannot contain an empty string.") raise ValueError("stop cannot contain an empty string.")
if self.stop and not self.detokenize: if self.stop and not self.detokenize:
raise ValueError( raise ValueError(
"stop strings are only supported when detokenize is True. " "stop strings are only supported when detokenize is True. "
"Set detokenize=True to use stop.") "Set detokenize=True to use stop.")
if self.best_of != self._real_n and self.output_kind == ( if self.best_of != self._real_n and self.output_kind == (
RequestOutputKind.DELTA): RequestOutputKind.DELTA):
raise ValueError("best_of must equal n to use output_kind=DELTA") raise ValueError("best_of must equal n to use output_kind=DELTA")
def _verify_greedy_sampling(self) -> None: def _verify_greedy_sampling(self) -> None:
if self.n > 1: if self.n > 1:
raise ValueError("n must be 1 when using greedy sampling, " raise ValueError("n must be 1 when using greedy sampling, "
f"got {self.n}.") f"got {self.n}.")
def update_from_generation_config( def update_from_generation_config(
self, self,
generation_config: Dict[str, Any], generation_config: Dict[str, Any],
model_eos_token_id: Optional[int] = None) -> None: model_eos_token_id: Optional[int] = None) -> None:
"""Update if there are non-default values from generation_config""" """Update if there are non-default values from generation_config"""
if model_eos_token_id is not None: if model_eos_token_id is not None:
# Add the eos token id into the sampling_params to support # Add the eos token id into the sampling_params to support
# min_tokens processing. # min_tokens processing.
self._all_stop_token_ids.add(model_eos_token_id) self._all_stop_token_ids.add(model_eos_token_id)
# Update eos_token_id for generation # Update eos_token_id for generation
if (eos_ids := generation_config.get("eos_token_id")) is not None: if (eos_ids := generation_config.get("eos_token_id")) is not None:
# it can be either int or list of int # it can be either int or list of int
eos_ids = {eos_ids} if isinstance(eos_ids, int) else set(eos_ids) eos_ids = {eos_ids} if isinstance(eos_ids, int) else set(eos_ids)
if model_eos_token_id is not None: if model_eos_token_id is not None:
# We don't need to include the primary eos_token_id in # We don't need to include the primary eos_token_id in
# stop_token_ids since it's handled separately for stopping # stop_token_ids since it's handled separately for stopping
# purposes. # purposes.
eos_ids.discard(model_eos_token_id) eos_ids.discard(model_eos_token_id)
if eos_ids: if eos_ids:
self._all_stop_token_ids.update(eos_ids) self._all_stop_token_ids.update(eos_ids)
if not self.ignore_eos: if not self.ignore_eos:
eos_ids.update(self.stop_token_ids) eos_ids.update(self.stop_token_ids)
self.stop_token_ids = list(eos_ids) self.stop_token_ids = list(eos_ids)
@cached_property @cached_property
def sampling_type(self) -> SamplingType: def sampling_type(self) -> SamplingType:
if self.temperature < _SAMPLING_EPS: if self.temperature < _SAMPLING_EPS:
return SamplingType.GREEDY return SamplingType.GREEDY
if self.seed is not None: if self.seed is not None:
return SamplingType.RANDOM_SEED return SamplingType.RANDOM_SEED
return SamplingType.RANDOM return SamplingType.RANDOM
@property @property
def all_stop_token_ids(self) -> Set[int]: def all_stop_token_ids(self) -> Set[int]:
return self._all_stop_token_ids return self._all_stop_token_ids
def clone(self) -> "SamplingParams": def clone(self) -> "SamplingParams":
"""Deep copy excluding LogitsProcessor objects. """Deep copy excluding LogitsProcessor objects.
LogitsProcessor objects are excluded because they may contain an LogitsProcessor objects are excluded because they may contain an
arbitrary, nontrivial amount of data. arbitrary, nontrivial amount of data.
See https://github.com/vllm-project/vllm/issues/3087 See https://github.com/vllm-project/vllm/issues/3087
""" """
logit_processor_refs = None if self.logits_processors is None else { logit_processor_refs = None if self.logits_processors is None else {
id(lp): lp id(lp): lp
for lp in self.logits_processors for lp in self.logits_processors
} }
return copy.deepcopy(self, memo=logit_processor_refs) return copy.deepcopy(self, memo=logit_processor_refs)
def __repr__(self) -> str: def __repr__(self) -> str:
return ( return (
f"SamplingParams(n={self.n}, " f"SamplingParams(n={self.n}, "
f"presence_penalty={self.presence_penalty}, " f"presence_penalty={self.presence_penalty}, "
f"frequency_penalty={self.frequency_penalty}, " f"frequency_penalty={self.frequency_penalty}, "
f"repetition_penalty={self.repetition_penalty}, " f"repetition_penalty={self.repetition_penalty}, "
f"temperature={self.temperature}, " f"temperature={self.temperature}, "
f"top_p={self.top_p}, " f"top_p={self.top_p}, "
f"top_k={self.top_k}, " f"top_k={self.top_k}, "
f"min_p={self.min_p}, " f"min_p={self.min_p}, "
f"seed={self.seed}, " f"seed={self.seed}, "
f"stop={self.stop}, " f"stop={self.stop}, "
f"stop_token_ids={self.stop_token_ids}, " f"stop_token_ids={self.stop_token_ids}, "
f"include_stop_str_in_output={self.include_stop_str_in_output}, " f"include_stop_str_in_output={self.include_stop_str_in_output}, "
f"ignore_eos={self.ignore_eos}, " f"ignore_eos={self.ignore_eos}, "
f"max_tokens={self.max_tokens}, " f"max_tokens={self.max_tokens}, "
f"min_tokens={self.min_tokens}, " f"min_tokens={self.min_tokens}, "
f"logprobs={self.logprobs}, " f"logprobs={self.logprobs}, "
f"prompt_logprobs={self.prompt_logprobs}, " f"prompt_logprobs={self.prompt_logprobs}, "
"prompt_logprob_positions=" "prompt_logprob_positions="
f"{self.prompt_logprob_positions}, " f"{self.prompt_logprob_positions}, "
f"skip_special_tokens={self.skip_special_tokens}, " f"skip_special_tokens={self.skip_special_tokens}, "
"spaces_between_special_tokens=" "spaces_between_special_tokens="
f"{self.spaces_between_special_tokens}, " f"{self.spaces_between_special_tokens}, "
f"truncate_prompt_tokens={self.truncate_prompt_tokens}), " f"truncate_prompt_tokens={self.truncate_prompt_tokens}), "
f"guided_decoding={self.guided_decoding}") f"guided_decoding={self.guided_decoding}")
class BeamSearchParams( class BeamSearchParams(
msgspec.Struct, msgspec.Struct,
omit_defaults=True, # type: ignore[call-arg] omit_defaults=True, # type: ignore[call-arg]
# required for @cached_property. # required for @cached_property.
dict=True): # type: ignore[call-arg] dict=True): # type: ignore[call-arg]
"""Beam search parameters for text generation.""" """Beam search parameters for text generation."""
beam_width: int beam_width: int
max_tokens: int max_tokens: int
ignore_eos: bool = False ignore_eos: bool = False
temperature: float = 0.0 temperature: float = 0.0
length_penalty: float = 1.0 length_penalty: float = 1.0