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:
10
.gitattributes
vendored
Normal file
10
.gitattributes
vendored
Normal 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
@@ -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)
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
@@ -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"]
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
@@ -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
@@ -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",
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user