diff --git a/qwen3_6_scripts/arg_utils.py b/qwen3_6_scripts/arg_utils.py
new file mode 100644
index 00000000..8faa8a0c
--- /dev/null
+++ b/qwen3_6_scripts/arg_utils.py
@@ -0,0 +1,1138 @@
+import argparse
+import dataclasses
+import json
+from dataclasses import dataclass
+from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Mapping, Optional,
+ Tuple, Type, Union)
+
+import torch
+
+import vllm.envs as envs
+from vllm.config import (CacheConfig, ConfigFormat, DecodingConfig,
+ DeviceConfig, EngineConfig, LoadConfig, LoadFormat,
+ LoRAConfig, ModelConfig, ObservabilityConfig,
+ ParallelConfig, PromptAdapterConfig, SchedulerConfig,
+ SpeculativeConfig, TokenizerPoolConfig)
+from vllm.executor.executor_base import ExecutorBase
+from vllm.logger import init_logger
+from vllm.model_executor.layers.quantization import QUANTIZATION_METHODS
+from vllm.transformers_utils.config import (
+ maybe_register_config_serialize_by_value)
+from vllm.transformers_utils.utils import check_gguf_file
+from vllm.utils import FlexibleArgumentParser
+
+if TYPE_CHECKING:
+ from vllm.transformers_utils.tokenizer_group import BaseTokenizerGroup
+
+logger = init_logger(__name__)
+
+ALLOWED_DETAILED_TRACE_MODULES = ["model", "worker", "all"]
+
+DEVICE_OPTIONS = [
+ "auto",
+ "cuda",
+ "neuron",
+ "cpu",
+ "openvino",
+ "tpu",
+ "xpu",
+]
+
+
+def nullable_str(val: str):
+ if not val or val == "None":
+ return None
+ return val
+
+
+def nullable_kvs(val: str) -> Optional[Mapping[str, int]]:
+ """Parses a string containing comma separate key [str] to value [int]
+ pairs into a dictionary.
+
+ Args:
+ val: String value to be parsed.
+
+ Returns:
+ Dictionary with parsed values.
+ """
+ if len(val) == 0:
+ return None
+
+ out_dict: Dict[str, int] = {}
+ for item in val.split(","):
+ kv_parts = [part.lower().strip() for part in item.split("=")]
+ if len(kv_parts) != 2:
+ raise argparse.ArgumentTypeError(
+ "Each item should be in the form KEY=VALUE")
+ key, value = kv_parts
+
+ try:
+ parsed_value = int(value)
+ except ValueError as exc:
+ msg = f"Failed to parse value of item {key}={value}"
+ raise argparse.ArgumentTypeError(msg) from exc
+
+ if key in out_dict and out_dict[key] != parsed_value:
+ raise argparse.ArgumentTypeError(
+ f"Conflicting values specified for key: {key}")
+ out_dict[key] = parsed_value
+
+ return out_dict
+
+
+@dataclass
+class EngineArgs:
+ """Arguments for vLLM engine."""
+ model: str = 'facebook/opt-125m'
+ served_model_name: Optional[Union[str, List[str]]] = None
+ tokenizer: Optional[str] = None
+ skip_tokenizer_init: bool = False
+ tokenizer_mode: str = 'auto'
+ trust_remote_code: bool = False
+ download_dir: Optional[str] = None
+ load_format: str = 'auto'
+ config_format: str = 'auto'
+ dtype: str = 'auto'
+ kv_cache_dtype: str = 'auto'
+ quantization_param_path: Optional[str] = None
+ seed: int = 0
+ max_model_len: Optional[int] = None
+ worker_use_ray: bool = False
+ # Note: Specifying a custom executor backend by passing a class
+ # is intended for expert use only. The API may change without
+ # notice.
+ distributed_executor_backend: Optional[Union[str,
+ Type[ExecutorBase]]] = None
+ pipeline_parallel_size: int = 1
+ tensor_parallel_size: int = 1
+ max_parallel_loading_workers: Optional[int] = None
+ block_size: int = 16
+ enable_prefix_caching: bool = False
+ disable_sliding_window: bool = False
+ use_v2_block_manager: bool = True
+ swap_space: float = 4 # GiB
+ cpu_offload_gb: float = 0 # GiB
+ gpu_memory_utilization: float = 0.90
+ max_num_batched_tokens: Optional[int] = None
+ max_num_seqs: int = 256
+ max_logprobs: int = 20 # Default value for OpenAI Chat Completions API
+ disable_log_stats: bool = False
+ revision: Optional[str] = None
+ code_revision: Optional[str] = None
+ rope_scaling: Optional[dict] = None
+ rope_theta: Optional[float] = None
+ tokenizer_revision: Optional[str] = None
+ quantization: Optional[str] = None
+ enforce_eager: Optional[bool] = None
+ max_context_len_to_capture: Optional[int] = None
+ max_seq_len_to_capture: int = 8192
+ disable_custom_all_reduce: bool = False
+ tokenizer_pool_size: int = 0
+ # Note: Specifying a tokenizer pool by passing a class
+ # is intended for expert use only. The API may change without
+ # notice.
+ tokenizer_pool_type: Union[str, Type["BaseTokenizerGroup"]] = "ray"
+ tokenizer_pool_extra_config: Optional[dict] = None
+ limit_mm_per_prompt: Optional[Mapping[str, int]] = None
+ enable_lora: bool = False
+ max_loras: int = 1
+ max_lora_rank: int = 16
+ enable_prompt_adapter: bool = False
+ max_prompt_adapters: int = 1
+ max_prompt_adapter_token: int = 0
+ fully_sharded_loras: bool = False
+ lora_extra_vocab_size: int = 256
+ long_lora_scaling_factors: Optional[Tuple[float]] = None
+ lora_dtype: Optional[Union[str, torch.dtype]] = 'auto'
+ max_cpu_loras: Optional[int] = None
+ device: str = 'auto'
+ num_scheduler_steps: int = 1
+ multi_step_stream_outputs: bool = True
+ ray_workers_use_nsight: bool = False
+ num_gpu_blocks_override: Optional[int] = None
+ num_lookahead_slots: int = 0
+ model_loader_extra_config: Optional[dict] = None
+ ignore_patterns: Optional[Union[str, List[str]]] = None
+ preemption_mode: Optional[str] = None
+
+ scheduler_delay_factor: float = 0.0
+ enable_chunked_prefill: Optional[bool] = None
+
+ guided_decoding_backend: str = 'outlines'
+ # Speculative decoding configuration.
+ speculative_model: Optional[str] = None
+ speculative_model_quantization: Optional[str] = None
+ speculative_draft_tensor_parallel_size: Optional[int] = None
+ num_speculative_tokens: Optional[int] = None
+ speculative_disable_mqa_scorer: Optional[bool] = False
+ speculative_max_model_len: Optional[int] = None
+ speculative_disable_by_batch_size: Optional[int] = None
+ ngram_prompt_lookup_max: Optional[int] = None
+ ngram_prompt_lookup_min: Optional[int] = None
+ spec_decoding_acceptance_method: str = 'rejection_sampler'
+ typical_acceptance_sampler_posterior_threshold: Optional[float] = None
+ typical_acceptance_sampler_posterior_alpha: Optional[float] = None
+ qlora_adapter_name_or_path: Optional[str] = None
+ disable_logprobs_during_spec_decoding: Optional[bool] = None
+
+ otlp_traces_endpoint: Optional[str] = None
+ collect_detailed_traces: Optional[str] = None
+ disable_async_output_proc: bool = False
+ override_neuron_config: Optional[Dict[str, Any]] = None
+ mm_processor_kwargs: Optional[Dict[str, Any]] = None
+ scheduling_policy: Literal["fcfs", "priority"] = "fcfs"
+
+ def __post_init__(self):
+ if self.tokenizer is None:
+ self.tokenizer = self.model
+
+ # Setup plugins
+ from vllm.plugins import load_general_plugins
+ load_general_plugins()
+
+ @staticmethod
+ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
+ """Shared CLI arguments for vLLM engine."""
+
+ # Model arguments
+ parser.add_argument(
+ '--model',
+ type=str,
+ default=EngineArgs.model,
+ help='Name or path of the huggingface model to use.')
+ parser.add_argument(
+ '--tokenizer',
+ type=nullable_str,
+ default=EngineArgs.tokenizer,
+ help='Name or path of the huggingface tokenizer to use. '
+ 'If unspecified, model name or path will be used.')
+ parser.add_argument(
+ '--skip-tokenizer-init',
+ action='store_true',
+ help='Skip initialization of tokenizer and detokenizer')
+ parser.add_argument(
+ '--revision',
+ type=nullable_str,
+ default=None,
+ help='The specific model version to use. It can be a branch '
+ 'name, a tag name, or a commit id. If unspecified, will use '
+ 'the default version.')
+ parser.add_argument(
+ '--code-revision',
+ type=nullable_str,
+ default=None,
+ help='The specific revision to use for the model code on '
+ 'Hugging Face Hub. It can be a branch name, a tag name, or a '
+ 'commit id. If unspecified, will use the default version.')
+ parser.add_argument(
+ '--tokenizer-revision',
+ type=nullable_str,
+ default=None,
+ help='Revision of the huggingface tokenizer to use. '
+ 'It can be a branch name, a tag name, or a commit id. '
+ 'If unspecified, will use the default version.')
+ parser.add_argument(
+ '--tokenizer-mode',
+ type=str,
+ default=EngineArgs.tokenizer_mode,
+ choices=['auto', 'slow', 'mistral'],
+ help='The tokenizer mode.\n\n* "auto" will use the '
+ 'fast tokenizer if available.\n* "slow" will '
+ 'always use the slow tokenizer. \n* '
+ '"mistral" will always use the `mistral_common` tokenizer.')
+ parser.add_argument('--trust-remote-code',
+ action='store_true',
+ help='Trust remote code from huggingface.')
+ parser.add_argument('--download-dir',
+ type=nullable_str,
+ default=EngineArgs.download_dir,
+ help='Directory to download and load the weights, '
+ 'default to the default cache dir of '
+ 'huggingface.')
+ parser.add_argument(
+ '--load-format',
+ type=str,
+ default=EngineArgs.load_format,
+ choices=[f.value for f in LoadFormat],
+ help='The format of the model weights to load.\n\n'
+ '* "auto" will try to load the weights in the safetensors format '
+ 'and fall back to the pytorch bin format if safetensors format '
+ 'is not available.\n'
+ '* "pt" will load the weights in the pytorch bin format.\n'
+ '* "safetensors" will load the weights in the safetensors format.\n'
+ '* "npcache" will load the weights in pytorch format and store '
+ 'a numpy cache to speed up the loading.\n'
+ '* "dummy" will initialize the weights with random values, '
+ 'which is mainly for profiling.\n'
+ '* "tensorizer" will load the weights using tensorizer from '
+ 'CoreWeave. See the Tensorize vLLM Model script in the Examples '
+ 'section for more information.\n'
+ '* "bitsandbytes" will load the weights using bitsandbytes '
+ 'quantization.\n')
+ parser.add_argument(
+ '--config-format',
+ default=EngineArgs.config_format,
+ choices=[f.value for f in ConfigFormat],
+ help='The format of the model config to load.\n\n'
+ '* "auto" will try to load the config in hf format '
+ 'if available else it will try to load in mistral format ')
+ parser.add_argument(
+ '--dtype',
+ type=str,
+ default=EngineArgs.dtype,
+ choices=[
+ 'auto', 'half', 'float16', 'bfloat16', 'float', 'float32'
+ ],
+ help='Data type for model weights and activations.\n\n'
+ '* "auto" will use FP16 precision for FP32 and FP16 models, and '
+ 'BF16 precision for BF16 models.\n'
+ '* "half" for FP16. Recommended for AWQ quantization.\n'
+ '* "float16" is the same as "half".\n'
+ '* "bfloat16" for a balance between precision and range.\n'
+ '* "float" is shorthand for FP32 precision.\n'
+ '* "float32" for FP32 precision.')
+ parser.add_argument(
+ '--kv-cache-dtype',
+ type=str,
+ choices=['auto', 'fp8', 'fp8_e5m2', 'fp8_e4m3'],
+ default=EngineArgs.kv_cache_dtype,
+ help='Data type for kv cache storage. If "auto", will use model '
+ 'data type. CUDA 11.8+ supports fp8 (=fp8_e4m3) and fp8_e5m2. '
+ 'ROCm (AMD GPU) supports fp8 (=fp8_e4m3)')
+ parser.add_argument(
+ '--quantization-param-path',
+ type=nullable_str,
+ default=None,
+ help='Path to the JSON file containing the KV cache '
+ 'scaling factors. This should generally be supplied, when '
+ 'KV cache dtype is FP8. Otherwise, KV cache scaling factors '
+ 'default to 1.0, which may cause accuracy issues. '
+ 'FP8_E5M2 (without scaling) is only supported on cuda version'
+ 'greater than 11.8. On ROCm (AMD GPU), FP8_E4M3 is instead '
+ 'supported for common inference criteria.')
+ parser.add_argument('--max-model-len',
+ type=int,
+ default=EngineArgs.max_model_len,
+ help='Model context length. If unspecified, will '
+ 'be automatically derived from the model config.')
+ parser.add_argument(
+ '--guided-decoding-backend',
+ type=str,
+ default='outlines',
+ choices=['outlines', 'lm-format-enforcer'],
+ help='Which engine will be used for guided decoding'
+ ' (JSON schema / regex etc) by default. Currently support '
+ 'https://github.com/outlines-dev/outlines and '
+ 'https://github.com/noamgat/lm-format-enforcer.'
+ ' Can be overridden per request via guided_decoding_backend'
+ ' parameter.')
+ # Parallel arguments
+ parser.add_argument(
+ '--distributed-executor-backend',
+ choices=['ray', 'mp'],
+ default=EngineArgs.distributed_executor_backend,
+ help='Backend to use for distributed serving. When more than 1 GPU '
+ 'is used, will be automatically set to "ray" if installed '
+ 'or "mp" (multiprocessing) otherwise.')
+ parser.add_argument(
+ '--worker-use-ray',
+ action='store_true',
+ help='Deprecated, use --distributed-executor-backend=ray.')
+ parser.add_argument('--pipeline-parallel-size',
+ '-pp',
+ type=int,
+ default=EngineArgs.pipeline_parallel_size,
+ help='Number of pipeline stages.')
+ parser.add_argument('--tensor-parallel-size',
+ '-tp',
+ type=int,
+ default=EngineArgs.tensor_parallel_size,
+ help='Number of tensor parallel replicas.')
+ parser.add_argument(
+ '--max-parallel-loading-workers',
+ type=int,
+ default=EngineArgs.max_parallel_loading_workers,
+ help='Load model sequentially in multiple batches, '
+ 'to avoid RAM OOM when using tensor '
+ 'parallel and large models.')
+ parser.add_argument(
+ '--ray-workers-use-nsight',
+ action='store_true',
+ help='If specified, use nsight to profile Ray workers.')
+ # KV cache arguments
+ parser.add_argument('--block-size',
+ type=int,
+ default=EngineArgs.block_size,
+ choices=[8, 16, 32],
+ help='Token block size for contiguous chunks of '
+ 'tokens. This is ignored on neuron devices and '
+ 'set to max-model-len')
+
+ parser.add_argument('--enable-prefix-caching',
+ action='store_true',
+ help='Enables automatic prefix caching.')
+ parser.add_argument('--disable-sliding-window',
+ action='store_true',
+ help='Disables sliding window, '
+ 'capping to sliding window size')
+ parser.add_argument(
+ '--use-v2-block-manager',
+ default=EngineArgs.use_v2_block_manager,
+ action='store_true',
+ help='Use BlockSpaceMangerV2. By default this is set to True. '
+ 'Set to False to use BlockSpaceManagerV1')
+ parser.add_argument(
+ '--num-lookahead-slots',
+ type=int,
+ default=EngineArgs.num_lookahead_slots,
+ help='Experimental scheduling config necessary for '
+ 'speculative decoding. This will be replaced by '
+ 'speculative config in the future; it is present '
+ 'to enable correctness tests until then.')
+
+ parser.add_argument('--seed',
+ type=int,
+ default=EngineArgs.seed,
+ help='Random seed for operations.')
+ parser.add_argument('--swap-space',
+ type=float,
+ default=EngineArgs.swap_space,
+ help='CPU swap space size (GiB) per GPU.')
+ parser.add_argument(
+ '--cpu-offload-gb',
+ type=float,
+ default=0,
+ help='The space in GiB to offload to CPU, per GPU. '
+ 'Default is 0, which means no offloading. Intuitively, '
+ 'this argument can be seen as a virtual way to increase '
+ 'the GPU memory size. For example, if you have one 24 GB '
+ 'GPU and set this to 10, virtually you can think of it as '
+ 'a 34 GB GPU. Then you can load a 13B model with BF16 weight,'
+ 'which requires at least 26GB GPU memory. Note that this '
+ 'requires fast CPU-GPU interconnect, as part of the model is'
+ 'loaded from CPU memory to GPU memory on the fly in each '
+ 'model forward pass.')
+ parser.add_argument(
+ '--gpu-memory-utilization',
+ type=float,
+ default=EngineArgs.gpu_memory_utilization,
+ help='The fraction of GPU memory to be used for the model '
+ 'executor, which can range from 0 to 1. For example, a value of '
+ '0.5 would imply 50%% GPU memory utilization. If unspecified, '
+ 'will use the default value of 0.9.')
+ parser.add_argument(
+ '--num-gpu-blocks-override',
+ type=int,
+ default=None,
+ help='If specified, ignore GPU profiling result and use this number'
+ 'of GPU blocks. Used for testing preemption.')
+ parser.add_argument('--max-num-batched-tokens',
+ type=int,
+ default=EngineArgs.max_num_batched_tokens,
+ help='Maximum number of batched tokens per '
+ 'iteration.')
+ parser.add_argument('--max-num-seqs',
+ type=int,
+ default=EngineArgs.max_num_seqs,
+ help='Maximum number of sequences per iteration.')
+ parser.add_argument(
+ '--max-logprobs',
+ type=int,
+ default=EngineArgs.max_logprobs,
+ help=('Max number of log probs to return logprobs is specified in'
+ ' SamplingParams.'))
+ parser.add_argument('--disable-log-stats',
+ action='store_true',
+ help='Disable logging statistics.')
+ # Quantization settings.
+ parser.add_argument('--quantization',
+ '-q',
+ type=nullable_str,
+ choices=[*QUANTIZATION_METHODS, None],
+ default=EngineArgs.quantization,
+ help='Method used to quantize the weights. If '
+ 'None, we first check the `quantization_config` '
+ 'attribute in the model config file. If that is '
+ 'None, we assume the model weights are not '
+ 'quantized and use `dtype` to determine the data '
+ 'type of the weights.')
+ parser.add_argument('--rope-scaling',
+ default=None,
+ type=json.loads,
+ help='RoPE scaling configuration in JSON format. '
+ 'For example, {"type":"dynamic","factor":2.0}')
+ parser.add_argument('--rope-theta',
+ default=None,
+ type=float,
+ help='RoPE theta. Use with `rope_scaling`. In '
+ 'some cases, changing the RoPE theta improves the '
+ 'performance of the scaled model.')
+ parser.add_argument('--enforce-eager',
+ action='store_true',
+ help='Always use eager-mode PyTorch. If False, '
+ 'will use eager mode and CUDA graph in hybrid '
+ 'for maximal performance and flexibility.')
+ parser.add_argument('--max-context-len-to-capture',
+ type=int,
+ default=EngineArgs.max_context_len_to_capture,
+ help='Maximum context length covered by CUDA '
+ 'graphs. When a sequence has context length '
+ 'larger than this, we fall back to eager mode. '
+ '(DEPRECATED. Use --max-seq-len-to-capture instead'
+ ')')
+ parser.add_argument('--max-seq-len-to-capture',
+ type=int,
+ default=EngineArgs.max_seq_len_to_capture,
+ help='Maximum sequence length covered by CUDA '
+ 'graphs. When a sequence has context length '
+ 'larger than this, we fall back to eager mode. '
+ 'Additionally for encoder-decoder models, if the '
+ 'sequence length of the encoder input is larger '
+ 'than this, we fall back to the eager mode.')
+ parser.add_argument('--disable-custom-all-reduce',
+ action='store_true',
+ default=EngineArgs.disable_custom_all_reduce,
+ help='See ParallelConfig.')
+ parser.add_argument('--tokenizer-pool-size',
+ type=int,
+ default=EngineArgs.tokenizer_pool_size,
+ help='Size of tokenizer pool to use for '
+ 'asynchronous tokenization. If 0, will '
+ 'use synchronous tokenization.')
+ parser.add_argument('--tokenizer-pool-type',
+ type=str,
+ default=EngineArgs.tokenizer_pool_type,
+ help='Type of tokenizer pool to use for '
+ 'asynchronous tokenization. Ignored '
+ 'if tokenizer_pool_size is 0.')
+ parser.add_argument('--tokenizer-pool-extra-config',
+ type=nullable_str,
+ default=EngineArgs.tokenizer_pool_extra_config,
+ help='Extra config for tokenizer pool. '
+ 'This should be a JSON string that will be '
+ 'parsed into a dictionary. Ignored if '
+ 'tokenizer_pool_size is 0.')
+
+ # Multimodal related configs
+ parser.add_argument(
+ '--limit-mm-per-prompt',
+ type=nullable_kvs,
+ default=EngineArgs.limit_mm_per_prompt,
+ # The default value is given in
+ # MultiModalRegistry.init_mm_limits_per_prompt
+ help=('For each multimodal plugin, limit how many '
+ 'input instances to allow for each prompt. '
+ 'Expects a comma-separated list of items, '
+ 'e.g.: `image=16,video=2` allows a maximum of 16 '
+ 'images and 2 videos per prompt. Defaults to 1 for '
+ 'each modality.'))
+ parser.add_argument(
+ '--mm-processor-kwargs',
+ default=None,
+ type=json.loads,
+ help=('Overrides for the multimodal input mapping/processing,'
+ 'e.g., image processor. For example: {"num_crops": 4}.'))
+
+ # LoRA related configs
+ parser.add_argument('--enable-lora',
+ action='store_true',
+ help='If True, enable handling of LoRA adapters.')
+ parser.add_argument('--max-loras',
+ type=int,
+ default=EngineArgs.max_loras,
+ help='Max number of LoRAs in a single batch.')
+ parser.add_argument('--max-lora-rank',
+ type=int,
+ default=EngineArgs.max_lora_rank,
+ help='Max LoRA rank.')
+ parser.add_argument(
+ '--lora-extra-vocab-size',
+ type=int,
+ default=EngineArgs.lora_extra_vocab_size,
+ help=('Maximum size of extra vocabulary that can be '
+ 'present in a LoRA adapter (added to the base '
+ 'model vocabulary).'))
+ parser.add_argument(
+ '--lora-dtype',
+ type=str,
+ default=EngineArgs.lora_dtype,
+ choices=['auto', 'float16', 'bfloat16', 'float32'],
+ help=('Data type for LoRA. If auto, will default to '
+ 'base model dtype.'))
+ parser.add_argument(
+ '--long-lora-scaling-factors',
+ type=nullable_str,
+ default=EngineArgs.long_lora_scaling_factors,
+ help=('Specify multiple scaling factors (which can '
+ 'be different from base model scaling factor '
+ '- see eg. Long LoRA) to allow for multiple '
+ 'LoRA adapters trained with those scaling '
+ 'factors to be used at the same time. If not '
+ 'specified, only adapters trained with the '
+ 'base model scaling factor are allowed.'))
+ parser.add_argument(
+ '--max-cpu-loras',
+ type=int,
+ default=EngineArgs.max_cpu_loras,
+ help=('Maximum number of LoRAs to store in CPU memory. '
+ 'Must be >= than max_num_seqs. '
+ 'Defaults to max_num_seqs.'))
+ parser.add_argument(
+ '--fully-sharded-loras',
+ action='store_true',
+ help=('By default, only half of the LoRA computation is '
+ 'sharded with tensor parallelism. '
+ 'Enabling this will use the fully sharded layers. '
+ 'At high sequence length, max rank or '
+ 'tensor parallel size, this is likely faster.'))
+ parser.add_argument('--enable-prompt-adapter',
+ action='store_true',
+ help='If True, enable handling of PromptAdapters.')
+ parser.add_argument('--max-prompt-adapters',
+ type=int,
+ default=EngineArgs.max_prompt_adapters,
+ help='Max number of PromptAdapters in a batch.')
+ parser.add_argument('--max-prompt-adapter-token',
+ type=int,
+ default=EngineArgs.max_prompt_adapter_token,
+ help='Max number of PromptAdapters tokens')
+ parser.add_argument("--device",
+ type=str,
+ default=EngineArgs.device,
+ choices=DEVICE_OPTIONS,
+ help='Device type for vLLM execution.')
+ parser.add_argument('--num-scheduler-steps',
+ type=int,
+ default=1,
+ help=('Maximum number of forward steps per '
+ 'scheduler call.'))
+
+ parser.add_argument(
+ '--multi-step-stream-outputs',
+ action=StoreBoolean,
+ default=EngineArgs.multi_step_stream_outputs,
+ nargs="?",
+ const="True",
+ help='If False, then multi-step will stream outputs at the end '
+ 'of all steps')
+ parser.add_argument(
+ '--scheduler-delay-factor',
+ type=float,
+ default=EngineArgs.scheduler_delay_factor,
+ help='Apply a delay (of delay factor multiplied by previous '
+ 'prompt latency) before scheduling next prompt.')
+ parser.add_argument(
+ '--enable-chunked-prefill',
+ action=StoreBoolean,
+ default=EngineArgs.enable_chunked_prefill,
+ nargs="?",
+ const="True",
+ help='If set, the prefill requests can be chunked based on the '
+ 'max_num_batched_tokens.')
+
+ parser.add_argument(
+ '--speculative-model',
+ type=nullable_str,
+ default=EngineArgs.speculative_model,
+ help=
+ 'The name of the draft model to be used in speculative decoding.')
+ # Quantization settings for speculative model.
+ parser.add_argument(
+ '--speculative-model-quantization',
+ type=nullable_str,
+ choices=[*QUANTIZATION_METHODS, None],
+ default=EngineArgs.speculative_model_quantization,
+ help='Method used to quantize the weights of speculative model. '
+ 'If None, we first check the `quantization_config` '
+ 'attribute in the model config file. If that is '
+ 'None, we assume the model weights are not '
+ 'quantized and use `dtype` to determine the data '
+ 'type of the weights.')
+ parser.add_argument(
+ '--num-speculative-tokens',
+ type=int,
+ default=EngineArgs.num_speculative_tokens,
+ help='The number of speculative tokens to sample from '
+ 'the draft model in speculative decoding.')
+ parser.add_argument(
+ '--speculative-disable-mqa-scorer',
+ action='store_true',
+ help=
+ 'If set to True, the MQA scorer will be disabled in speculative '
+ ' and fall back to batch expansion')
+ parser.add_argument(
+ '--speculative-draft-tensor-parallel-size',
+ '-spec-draft-tp',
+ type=int,
+ default=EngineArgs.speculative_draft_tensor_parallel_size,
+ help='Number of tensor parallel replicas for '
+ 'the draft model in speculative decoding.')
+
+ parser.add_argument(
+ '--speculative-max-model-len',
+ type=int,
+ default=EngineArgs.speculative_max_model_len,
+ help='The maximum sequence length supported by the '
+ 'draft model. Sequences over this length will skip '
+ 'speculation.')
+
+ parser.add_argument(
+ '--speculative-disable-by-batch-size',
+ type=int,
+ default=EngineArgs.speculative_disable_by_batch_size,
+ help='Disable speculative decoding for new incoming requests '
+ 'if the number of enqueue requests is larger than this value.')
+
+ parser.add_argument(
+ '--ngram-prompt-lookup-max',
+ type=int,
+ default=EngineArgs.ngram_prompt_lookup_max,
+ help='Max size of window for ngram prompt lookup in speculative '
+ 'decoding.')
+
+ parser.add_argument(
+ '--ngram-prompt-lookup-min',
+ type=int,
+ default=EngineArgs.ngram_prompt_lookup_min,
+ help='Min size of window for ngram prompt lookup in speculative '
+ 'decoding.')
+
+ parser.add_argument(
+ '--spec-decoding-acceptance-method',
+ type=str,
+ default=EngineArgs.spec_decoding_acceptance_method,
+ choices=['rejection_sampler', 'typical_acceptance_sampler'],
+ help='Specify the acceptance method to use during draft token '
+ 'verification in speculative decoding. Two types of acceptance '
+ 'routines are supported: '
+ '1) RejectionSampler which does not allow changing the '
+ 'acceptance rate of draft tokens, '
+ '2) TypicalAcceptanceSampler which is configurable, allowing for '
+ 'a higher acceptance rate at the cost of lower quality, '
+ 'and vice versa.')
+
+ parser.add_argument(
+ '--typical-acceptance-sampler-posterior-threshold',
+ type=float,
+ default=EngineArgs.typical_acceptance_sampler_posterior_threshold,
+ help='Set the lower bound threshold for the posterior '
+ 'probability of a token to be accepted. This threshold is '
+ 'used by the TypicalAcceptanceSampler to make sampling decisions '
+ 'during speculative decoding. Defaults to 0.09')
+
+ parser.add_argument(
+ '--typical-acceptance-sampler-posterior-alpha',
+ type=float,
+ default=EngineArgs.typical_acceptance_sampler_posterior_alpha,
+ help='A scaling factor for the entropy-based threshold for token '
+ 'acceptance in the TypicalAcceptanceSampler. Typically defaults '
+ 'to sqrt of --typical-acceptance-sampler-posterior-threshold '
+ 'i.e. 0.3')
+
+ parser.add_argument(
+ '--disable-logprobs-during-spec-decoding',
+ action=StoreBoolean,
+ default=EngineArgs.disable_logprobs_during_spec_decoding,
+ nargs="?",
+ const="True",
+ help='If set to True, token log probabilities are not returned '
+ 'during speculative decoding. If set to False, log probabilities '
+ 'are returned according to the settings in SamplingParams. If '
+ 'not specified, it defaults to True. Disabling log probabilities '
+ 'during speculative decoding reduces latency by skipping logprob '
+ 'calculation in proposal sampling, target sampling, and after '
+ 'accepted tokens are determined.')
+
+ parser.add_argument('--model-loader-extra-config',
+ type=nullable_str,
+ default=EngineArgs.model_loader_extra_config,
+ help='Extra config for model loader. '
+ 'This will be passed to the model loader '
+ 'corresponding to the chosen load_format. '
+ 'This should be a JSON string that will be '
+ 'parsed into a dictionary.')
+ parser.add_argument(
+ '--ignore-patterns',
+ action="append",
+ type=str,
+ default=[],
+ help="The pattern(s) to ignore when loading the model."
+ "Default to 'original/**/*' to avoid repeated loading of llama's "
+ "checkpoints.")
+ parser.add_argument(
+ '--preemption-mode',
+ type=str,
+ default=None,
+ help='If \'recompute\', the engine performs preemption by '
+ 'recomputing; If \'swap\', the engine performs preemption by '
+ 'block swapping.')
+
+ parser.add_argument(
+ "--served-model-name",
+ nargs="+",
+ type=str,
+ default=None,
+ help="The model name(s) used in the API. If multiple "
+ "names are provided, the server will respond to any "
+ "of the provided names. The model name in the model "
+ "field of a response will be the first name in this "
+ "list. If not specified, the model name will be the "
+ "same as the `--model` argument. Noted that this name(s)"
+ "will also be used in `model_name` tag content of "
+ "prometheus metrics, if multiple names provided, metrics"
+ "tag will take the first one.")
+ parser.add_argument('--qlora-adapter-name-or-path',
+ type=str,
+ default=None,
+ help='Name or path of the QLoRA adapter.')
+
+ parser.add_argument(
+ '--otlp-traces-endpoint',
+ type=str,
+ default=None,
+ help='Target URL to which OpenTelemetry traces will be sent.')
+ parser.add_argument(
+ '--collect-detailed-traces',
+ type=str,
+ default=None,
+ help="Valid choices are " +
+ ",".join(ALLOWED_DETAILED_TRACE_MODULES) +
+ ". It makes sense to set this only if --otlp-traces-endpoint is"
+ " set. If set, it will collect detailed traces for the specified "
+ "modules. This involves use of possibly costly and or blocking "
+ "operations and hence might have a performance impact.")
+
+ parser.add_argument(
+ '--disable-async-output-proc',
+ action='store_true',
+ default=EngineArgs.disable_async_output_proc,
+ help="Disable async output processing. This may result in "
+ "lower performance.")
+ parser.add_argument(
+ '--override-neuron-config',
+ type=json.loads,
+ default=None,
+ help="Override or set neuron device configuration. "
+ "e.g. {\"cast_logits_dtype\": \"bloat16\"}.'")
+
+ parser.add_argument(
+ '--scheduling-policy',
+ choices=['fcfs', 'priority'],
+ default="fcfs",
+ help='The scheduling policy to use. "fcfs" (first come first served'
+ ', i.e. requests are handled in order of arrival; default) '
+ 'or "priority" (requests are handled based on given '
+ 'priority (lower value means earlier handling) and time of '
+ 'arrival deciding any ties).')
+
+ return parser
+
+ @classmethod
+ def from_cli_args(cls, args: argparse.Namespace):
+ # Get the list of attributes of this dataclass.
+ attrs = [attr.name for attr in dataclasses.fields(cls)]
+ # Set the attributes from the parsed arguments.
+ engine_args = cls(**{attr: getattr(args, attr) for attr in attrs})
+ return engine_args
+
+ def create_model_config(self) -> ModelConfig:
+ return ModelConfig(
+ model=self.model,
+ tokenizer=self.tokenizer,
+ tokenizer_mode=self.tokenizer_mode,
+ trust_remote_code=self.trust_remote_code,
+ dtype=self.dtype,
+ seed=self.seed,
+ revision=self.revision,
+ code_revision=self.code_revision,
+ rope_scaling=self.rope_scaling,
+ rope_theta=self.rope_theta,
+ tokenizer_revision=self.tokenizer_revision,
+ max_model_len=self.max_model_len,
+ quantization=self.quantization,
+ quantization_param_path=self.quantization_param_path,
+ enforce_eager=True,
+ max_context_len_to_capture=self.max_context_len_to_capture,
+ max_seq_len_to_capture=self.max_seq_len_to_capture,
+ max_logprobs=self.max_logprobs,
+ disable_sliding_window=self.disable_sliding_window,
+ skip_tokenizer_init=self.skip_tokenizer_init,
+ served_model_name=self.served_model_name,
+ limit_mm_per_prompt=self.limit_mm_per_prompt,
+ use_async_output_proc=not self.disable_async_output_proc,
+ override_neuron_config=self.override_neuron_config,
+ config_format=self.config_format,
+ mm_processor_kwargs=self.mm_processor_kwargs,
+ )
+
+ def create_load_config(self) -> LoadConfig:
+ return LoadConfig(
+ load_format=self.load_format,
+ download_dir=self.download_dir,
+ model_loader_extra_config=self.model_loader_extra_config,
+ ignore_patterns=self.ignore_patterns,
+ )
+
+ def create_engine_config(self) -> EngineConfig:
+ # gguf file needs a specific model loader and doesn't use hf_repo
+ if check_gguf_file(self.model):
+ self.quantization = self.load_format = "gguf"
+
+ # bitsandbytes quantization needs a specific model loader
+ # so we make sure the quant method and the load format are consistent
+ if (self.quantization == "bitsandbytes" or
+ self.qlora_adapter_name_or_path is not None) and \
+ self.load_format != "bitsandbytes":
+ raise ValueError(
+ "BitsAndBytes quantization and QLoRA adapter only support "
+ f"'bitsandbytes' load format, but got {self.load_format}")
+
+ if (self.load_format == "bitsandbytes" or
+ self.qlora_adapter_name_or_path is not None) and \
+ self.quantization != "bitsandbytes":
+ raise ValueError(
+ "BitsAndBytes load format and QLoRA adapter only support "
+ f"'bitsandbytes' quantization, but got {self.quantization}")
+
+ assert self.cpu_offload_gb >= 0, (
+ "CPU offload space must be non-negative"
+ f", but got {self.cpu_offload_gb}")
+
+ device_config = DeviceConfig(device=self.device)
+ model_config = self.create_model_config()
+
+ if model_config.is_multimodal_model:
+ if self.enable_prefix_caching:
+ logger.warning(
+ "--enable-prefix-caching is currently not "
+ "supported for multimodal models and has been disabled.")
+ self.enable_prefix_caching = False
+
+ maybe_register_config_serialize_by_value(self.trust_remote_code)
+
+ cache_config = CacheConfig(
+ block_size=self.block_size if self.device != "neuron" else
+ self.max_model_len, # neuron needs block_size = max_model_len
+ gpu_memory_utilization=self.gpu_memory_utilization,
+ swap_space=self.swap_space,
+ cache_dtype=self.kv_cache_dtype,
+ is_attention_free=model_config.is_attention_free,
+ num_gpu_blocks_override=self.num_gpu_blocks_override,
+ sliding_window=model_config.get_sliding_window(),
+ enable_prefix_caching=self.enable_prefix_caching,
+ cpu_offload_gb=self.cpu_offload_gb,
+ )
+ parallel_config = ParallelConfig(
+ pipeline_parallel_size=self.pipeline_parallel_size,
+ tensor_parallel_size=self.tensor_parallel_size,
+ worker_use_ray=self.worker_use_ray,
+ max_parallel_loading_workers=self.max_parallel_loading_workers,
+ disable_custom_all_reduce=True,
+ tokenizer_pool_config=TokenizerPoolConfig.create_config(
+ self.tokenizer_pool_size,
+ self.tokenizer_pool_type,
+ self.tokenizer_pool_extra_config,
+ ),
+ ray_workers_use_nsight=self.ray_workers_use_nsight,
+ distributed_executor_backend=self.distributed_executor_backend)
+
+ max_model_len = model_config.max_model_len
+ use_long_context = max_model_len > 32768
+
+ if self.enable_chunked_prefill is None:
+ # If not explicitly set, enable chunked prefill by default for
+ # long context (> 32K) models. This is to avoid OOM errors in the
+ # initial memory profiling phase.
+
+ # Chunked prefill is currently disabled for multimodal models by
+ # default.
+ if use_long_context and not model_config.is_multimodal_model:
+ is_gpu = device_config.device_type == "cuda"
+ use_sliding_window = (model_config.get_sliding_window()
+ is not None)
+ use_spec_decode = self.speculative_model is not None
+ if (is_gpu and not use_sliding_window and not use_spec_decode
+ and not self.enable_lora
+ and not self.enable_prompt_adapter):
+ pass # skip auto-enable: Q-tiling in _run_sdpa_fallback
+ # handles long-context memory without chunked prefill
+ if self.enable_chunked_prefill is None:
+ self.enable_chunked_prefill = False
+
+ if not self.enable_chunked_prefill and use_long_context:
+ logger.warning(
+ "The model has a long context length (%s). This may cause OOM "
+ "errors during the initial memory profiling phase, or result "
+ "in low performance due to small KV cache space. Consider "
+ "setting --max-model-len to a smaller value.", max_model_len)
+
+ if self.num_scheduler_steps > 1 and not self.use_v2_block_manager:
+ self.use_v2_block_manager = True
+ logger.warning(
+ "Enabled BlockSpaceManagerV2 because it is "
+ "required for multi-step (--num-scheduler-steps > 1)")
+
+ speculative_config = SpeculativeConfig.maybe_create_spec_config(
+ target_model_config=model_config,
+ target_parallel_config=parallel_config,
+ target_dtype=self.dtype,
+ speculative_model=self.speculative_model,
+ speculative_model_quantization = \
+ self.speculative_model_quantization,
+ speculative_draft_tensor_parallel_size = \
+ self.speculative_draft_tensor_parallel_size,
+ num_speculative_tokens=self.num_speculative_tokens,
+ speculative_disable_mqa_scorer=self.speculative_disable_mqa_scorer,
+ speculative_disable_by_batch_size=self.
+ speculative_disable_by_batch_size,
+ speculative_max_model_len=self.speculative_max_model_len,
+ enable_chunked_prefill=self.enable_chunked_prefill,
+ use_v2_block_manager=self.use_v2_block_manager,
+ disable_log_stats=self.disable_log_stats,
+ ngram_prompt_lookup_max=self.ngram_prompt_lookup_max,
+ ngram_prompt_lookup_min=self.ngram_prompt_lookup_min,
+ draft_token_acceptance_method=\
+ self.spec_decoding_acceptance_method,
+ typical_acceptance_sampler_posterior_threshold=self.
+ typical_acceptance_sampler_posterior_threshold,
+ typical_acceptance_sampler_posterior_alpha=self.
+ typical_acceptance_sampler_posterior_alpha,
+ disable_logprobs=self.disable_logprobs_during_spec_decoding,
+ )
+
+ # Reminder: Please update docs/source/serving/compatibility_matrix.rst
+ # If the feature combo become valid
+ if self.num_scheduler_steps > 1:
+ if speculative_config is not None:
+ raise ValueError("Speculative decoding is not supported with "
+ "multi-step (--num-scheduler-steps > 1)")
+ if self.enable_chunked_prefill and self.pipeline_parallel_size > 1:
+ raise ValueError("Multi-Step Chunked-Prefill is not supported "
+ "for pipeline-parallel-size > 1")
+
+ # make sure num_lookahead_slots is set the higher value depending on
+ # if we are using speculative decoding or multi-step
+ num_lookahead_slots = max(self.num_lookahead_slots,
+ self.num_scheduler_steps - 1)
+ num_lookahead_slots = num_lookahead_slots \
+ if speculative_config is None \
+ else speculative_config.num_lookahead_slots
+
+ scheduler_config = SchedulerConfig(
+ max_num_batched_tokens=self.max_num_batched_tokens,
+ max_num_seqs=self.max_num_seqs,
+ max_model_len=model_config.max_model_len,
+ use_v2_block_manager=self.use_v2_block_manager,
+ num_lookahead_slots=num_lookahead_slots,
+ delay_factor=self.scheduler_delay_factor,
+ enable_chunked_prefill=self.enable_chunked_prefill,
+ embedding_mode=model_config.embedding_mode,
+ is_multimodal_model=model_config.is_multimodal_model,
+ preemption_mode=self.preemption_mode,
+ num_scheduler_steps=self.num_scheduler_steps,
+ multi_step_stream_outputs=self.multi_step_stream_outputs,
+ send_delta_data=(envs.VLLM_USE_RAY_SPMD_WORKER
+ and parallel_config.use_ray),
+ policy=self.scheduling_policy,
+ )
+ lora_config = LoRAConfig(
+ max_lora_rank=self.max_lora_rank,
+ max_loras=self.max_loras,
+ fully_sharded_loras=self.fully_sharded_loras,
+ lora_extra_vocab_size=self.lora_extra_vocab_size,
+ long_lora_scaling_factors=self.long_lora_scaling_factors,
+ lora_dtype=self.lora_dtype,
+ max_cpu_loras=self.max_cpu_loras if self.max_cpu_loras
+ and self.max_cpu_loras > 0 else None) if self.enable_lora else None
+
+ if self.qlora_adapter_name_or_path is not None and \
+ self.qlora_adapter_name_or_path != "":
+ if self.model_loader_extra_config is None:
+ self.model_loader_extra_config = {}
+ self.model_loader_extra_config[
+ "qlora_adapter_name_or_path"] = self.qlora_adapter_name_or_path
+
+ load_config = self.create_load_config()
+
+ prompt_adapter_config = PromptAdapterConfig(
+ max_prompt_adapters=self.max_prompt_adapters,
+ max_prompt_adapter_token=self.max_prompt_adapter_token) \
+ if self.enable_prompt_adapter else None
+
+ decoding_config = DecodingConfig(
+ guided_decoding_backend=self.guided_decoding_backend)
+
+ detailed_trace_modules = []
+ if self.collect_detailed_traces is not None:
+ detailed_trace_modules = self.collect_detailed_traces.split(",")
+ for m in detailed_trace_modules:
+ if m not in ALLOWED_DETAILED_TRACE_MODULES:
+ raise ValueError(
+ f"Invalid module {m} in collect_detailed_traces. "
+ f"Valid modules are {ALLOWED_DETAILED_TRACE_MODULES}")
+ observability_config = ObservabilityConfig(
+ otlp_traces_endpoint=self.otlp_traces_endpoint,
+ collect_model_forward_time="model" in detailed_trace_modules
+ or "all" in detailed_trace_modules,
+ collect_model_execute_time="worker" in detailed_trace_modules
+ or "all" in detailed_trace_modules,
+ )
+
+ if (model_config.get_sliding_window() is not None
+ and scheduler_config.chunked_prefill_enabled
+ and not scheduler_config.use_v2_block_manager):
+ raise ValueError(
+ "Chunked prefill is not supported with sliding window. "
+ "Set --disable-sliding-window to disable sliding window.")
+
+ return EngineConfig(
+ model_config=model_config,
+ cache_config=cache_config,
+ parallel_config=parallel_config,
+ scheduler_config=scheduler_config,
+ device_config=device_config,
+ lora_config=lora_config,
+ speculative_config=speculative_config,
+ load_config=load_config,
+ decoding_config=decoding_config,
+ observability_config=observability_config,
+ prompt_adapter_config=prompt_adapter_config,
+ )
+
+
+@dataclass
+class AsyncEngineArgs(EngineArgs):
+ """Arguments for asynchronous vLLM engine."""
+ disable_log_requests: bool = False
+
+ @staticmethod
+ def add_cli_args(parser: FlexibleArgumentParser,
+ async_args_only: bool = False) -> FlexibleArgumentParser:
+ if not async_args_only:
+ parser = EngineArgs.add_cli_args(parser)
+ parser.add_argument('--disable-log-requests',
+ action='store_true',
+ help='Disable logging requests.')
+ return parser
+
+
+class StoreBoolean(argparse.Action):
+
+ def __call__(self, parser, namespace, values, option_string=None):
+ if values.lower() == "true":
+ setattr(namespace, self.dest, True)
+ elif values.lower() == "false":
+ setattr(namespace, self.dest, False)
+ else:
+ raise ValueError(f"Invalid boolean value: {values}. "
+ "Expected 'true' or 'false'.")
+
+
+# These functions are used by sphinx to build the documentation
+def _engine_args_parser():
+ return EngineArgs.add_cli_args(FlexibleArgumentParser())
+
+
+def _async_engine_args_parser():
+ return AsyncEngineArgs.add_cli_args(FlexibleArgumentParser(),
+ async_args_only=True)
diff --git a/qwen3_6_scripts/logits_processor.py b/qwen3_6_scripts/logits_processor.py
new file mode 100644
index 00000000..f0b0fd7d
--- /dev/null
+++ b/qwen3_6_scripts/logits_processor.py
@@ -0,0 +1,158 @@
+"""A layer that compute logits from hidden_stats."""
+import inspect
+from typing import Optional
+
+import torch
+import torch.nn as nn
+
+from vllm.distributed import (tensor_model_parallel_all_gather,
+ tensor_model_parallel_gather)
+from vllm.model_executor.layers.vocab_parallel_embedding import (
+ VocabParallelEmbedding)
+from vllm.model_executor.sampling_metadata import SamplingMetadata
+from vllm.platforms import current_platform
+
+
+class LogitsProcessor(nn.Module):
+ """Process logits and apply logits processors from sampling metadata.
+
+ This layer does the following:
+ 1. Gather logits from model hidden_states.
+ 2. Scale logits if needed.
+ 3. Apply logits processors (if any).
+ """
+
+ def __init__(self,
+ vocab_size: int,
+ org_vocab_size: Optional[int] = None,
+ scale: float = 1.0,
+ logits_as_input: bool = False,
+ soft_cap: Optional[float] = None) -> None:
+ """
+ Args:
+ scale: A scaling factor to apply to the logits.
+ """
+ super().__init__()
+ self.scale = scale
+ self.vocab_size = vocab_size
+ # Whether the input is logits (default is hidden states).
+ self.logits_as_input = logits_as_input
+ # original vocabulary size (without LoRA).
+ self.org_vocab_size = org_vocab_size or vocab_size
+ # Soft cap the logits. Used in Gemma 2.
+ self.soft_cap = soft_cap
+ # Whether to use gather or all-gather to gather the logits.
+ self.use_gather = not current_platform.is_tpu()
+
+ def forward(
+ self,
+ lm_head: VocabParallelEmbedding,
+ hidden_states: torch.Tensor,
+ sampling_metadata: SamplingMetadata,
+ embedding_bias: Optional[torch.Tensor] = None,
+ ) -> Optional[torch.Tensor]:
+ if self.logits_as_input:
+ logits = hidden_states
+ else:
+ hidden_states = _prune_hidden_states(hidden_states,
+ sampling_metadata)
+
+ # Get the logits for the next tokens.
+ if hidden_states.shape[0] > 0:
+ logits = self._get_logits(hidden_states, lm_head, embedding_bias)
+ else:
+ logits = torch.empty([0, lm_head.weight.shape[0]], device=hidden_states.device, dtype=hidden_states.dtype)
+ if logits is not None:
+ if self.soft_cap is not None:
+ logits = logits / self.soft_cap
+ logits = torch.tanh(logits)
+ logits = logits * self.soft_cap
+
+ if self.scale != 1.0:
+ logits *= self.scale
+
+ # Apply logits processors (if any).
+ logits = _apply_logits_processors(logits, sampling_metadata)
+
+ return logits
+
+ def _get_logits(
+ self,
+ hidden_states: torch.Tensor,
+ lm_head: VocabParallelEmbedding,
+ embedding_bias: Optional[torch.Tensor],
+ ) -> Optional[torch.Tensor]:
+ # Get the logits for the next tokens.
+ logits = lm_head.linear_method.apply(lm_head,
+ hidden_states,
+ bias=embedding_bias)
+ if self.use_gather:
+ # None may be returned for rank > 0
+ logits = tensor_model_parallel_gather(logits)
+ else:
+ # Gather is not supported for some devices such as TPUs.
+ # Use all-gather instead.
+ # NOTE(woosuk): Here, the outputs of every device should not be None
+ # because XLA requires strict SPMD among all devices. Every device
+ # should execute the same operations after gathering the logits.
+ logits = tensor_model_parallel_all_gather(logits)
+ # Remove paddings in vocab (if any).
+ if logits is not None:
+ logits = logits[..., :self.org_vocab_size]
+ return logits
+
+ def extra_repr(self) -> str:
+ s = f"vocab_size={self.vocab_size}"
+ s += f", forg_vocab_size={self.org_vocab_size}"
+ s += f", scale={self.scale}, logits_as_input={self.logits_as_input}"
+ return s
+
+
+def _prune_hidden_states(
+ hidden_states: torch.Tensor,
+ sampling_metadata: SamplingMetadata,
+) -> torch.Tensor:
+ return hidden_states.index_select(0,
+ sampling_metadata.selected_token_indices)
+
+
+def _apply_logits_processors(
+ logits: torch.Tensor,
+ sampling_metadata: SamplingMetadata,
+) -> torch.Tensor:
+ if sampling_metadata.seq_groups is None: # intermediate chunked-prefill chunk
+ return logits
+ found_logits_processors = False
+ logits_processed = 0
+ for seq_group in sampling_metadata.seq_groups:
+ seq_ids = seq_group.seq_ids
+ sampling_params = seq_group.sampling_params
+ logits_processors = sampling_params.logits_processors
+ if logits_processors:
+ found_logits_processors = True
+
+ for seq_id, logits_row_idx in zip(seq_ids,
+ seq_group.sample_indices):
+ logits_row = logits[logits_row_idx]
+ past_tokens_ids = seq_group.seq_data[seq_id].output_token_ids
+ prompt_tokens_ids = seq_group.seq_data[seq_id].prompt_token_ids
+
+ for logits_processor in logits_processors:
+ parameters = inspect.signature(logits_processor).parameters
+ if len(parameters) == 3:
+ logits_row = logits_processor(prompt_tokens_ids,
+ past_tokens_ids,
+ logits_row)
+ else:
+ logits_row = logits_processor(past_tokens_ids,
+ logits_row)
+
+ logits[logits_row_idx] = logits_row
+
+ logits_processed += len(seq_group.sample_indices) + len(
+ seq_group.prompt_logprob_indices)
+
+ if found_logits_processors:
+ # verifies that no rows in logits were missed unexpectedly
+ assert logits_processed == logits.shape[0]
+ return logits
diff --git a/qwen3_6_scripts/model_runner.py b/qwen3_6_scripts/model_runner.py
new file mode 100644
index 00000000..3c4f65a3
--- /dev/null
+++ b/qwen3_6_scripts/model_runner.py
@@ -0,0 +1,1937 @@
+import dataclasses
+import gc
+import inspect
+import itertools
+import time
+import warnings
+import weakref
+from dataclasses import dataclass
+from typing import (TYPE_CHECKING, Any, Callable, Dict, List, Optional, Set,
+ Tuple, Type, TypeVar, Union)
+
+import numpy as np
+import torch
+import torch.distributed
+import torch.nn as nn
+
+import vllm.envs as envs
+from vllm.attention import AttentionMetadata, get_attn_backend
+from vllm.attention.backends.abstract import AttentionState
+from vllm.attention.backends.utils import CommonAttentionState
+from vllm.compilation.compile_context import set_compile_context
+from vllm.compilation.levels import CompilationLevel
+from vllm.config import (CacheConfig, DeviceConfig, LoadConfig, LoRAConfig,
+ ModelConfig, ObservabilityConfig, ParallelConfig,
+ PromptAdapterConfig, SchedulerConfig)
+from vllm.core.scheduler import SchedulerOutputs
+from vllm.distributed import get_pp_group
+from vllm.distributed.parallel_state import graph_capture
+from vllm.forward_context import set_forward_context
+from vllm.inputs import INPUT_REGISTRY, InputRegistry
+from vllm.logger import init_logger
+from vllm.lora.layers import LoRAMapping
+from vllm.lora.request import LoRARequest
+from vllm.lora.worker_manager import LRUCacheWorkerLoRAManager
+from vllm.model_executor import SamplingMetadata, SamplingMetadataCache
+from vllm.model_executor.layers.rotary_embedding import MRotaryEmbedding
+from vllm.model_executor.layers.sampler import SamplerOutput
+from vllm.model_executor.model_loader import get_model
+from vllm.model_executor.model_loader.tensorizer import TensorizerConfig
+from vllm.model_executor.models import supports_lora, supports_multimodal
+from vllm.model_executor.models.utils import set_cpu_offload_max_bytes
+from vllm.multimodal import (MULTIMODAL_REGISTRY, BatchedTensorInputs,
+ MultiModalInputs, MultiModalRegistry)
+from vllm.prompt_adapter.layers import PromptAdapterMapping
+from vllm.prompt_adapter.request import PromptAdapterRequest
+from vllm.prompt_adapter.worker_manager import (
+ LRUCacheWorkerPromptAdapterManager)
+from vllm.sampling_params import SamplingParams
+from vllm.sequence import IntermediateTensors, SequenceGroupMetadata
+from vllm.utils import (DeviceMemoryProfiler, PyObjectCache, async_tensor_h2d,
+ flatten_2d_lists, is_hip, is_pin_memory_available,
+ supports_dynamo)
+from vllm.worker.model_runner_base import (
+ ModelRunnerBase, ModelRunnerInputBase, ModelRunnerInputBuilderBase,
+ _add_attn_metadata_broadcastable_dict,
+ _add_sampling_metadata_broadcastable_dict,
+ _init_attn_metadata_from_tensor_dict,
+ _init_sampling_metadata_from_tensor_dict, dump_input_when_exception)
+
+if TYPE_CHECKING:
+ from vllm.attention.backends.abstract import AttentionBackend
+
+logger = init_logger(__name__)
+
+LORA_WARMUP_RANK = 8
+_BATCH_SIZE_ALIGNMENT = 8
+# all the token sizes that **can** be captured by cudagraph.
+# they can be arbitrarily large.
+# currently it includes: 1, 2, 4, 8, 16, 24, 32, 40, ..., 8192.
+# the actual sizes to capture will be determined by the model,
+# depending on the model's max_num_seqs.
+# NOTE: _get_graph_batch_size needs to be updated if this list is changed.
+_BATCH_SIZES_TO_CAPTURE = [1, 2, 4] + [
+ _BATCH_SIZE_ALIGNMENT * i for i in range(1, 1025)
+]
+_NUM_WARMUP_ITERS = 2
+
+TModelInputForGPU = TypeVar('TModelInputForGPU', bound="ModelInputForGPU")
+
+# For now, bump up cache limits for recompilations during CUDA graph warmups.
+# torch._dynamo.config.cache_size_limit = 128
+# torch._dynamo.config.accumulated_cache_size_limit = 128
+
+
+@dataclass(frozen=True)
+class ModelInputForGPU(ModelRunnerInputBase):
+ """
+ This base class contains metadata needed for the base model forward pass
+ but not metadata for possible additional steps, e.g., sampling. Model
+ runners that run additional steps should subclass this method to add
+ additional fields.
+ """
+ input_tokens: Optional[torch.Tensor] = None
+ input_positions: Optional[torch.Tensor] = None
+ seq_lens: Optional[List[int]] = None
+ query_lens: Optional[List[int]] = None
+ lora_mapping: Optional["LoRAMapping"] = None
+ lora_requests: Optional[Set[LoRARequest]] = None
+ attn_metadata: Optional["AttentionMetadata"] = None
+ prompt_adapter_mapping: Optional[PromptAdapterMapping] = None
+ prompt_adapter_requests: Optional[Set[PromptAdapterRequest]] = None
+ multi_modal_kwargs: Optional[BatchedTensorInputs] = None
+ request_ids_to_seq_ids: Optional[Dict[str, List[int]]] = None
+ finished_requests_ids: Optional[List[str]] = None
+ virtual_engine: int = 0
+ async_callback: Optional[Callable] = None
+ seq_group_metadata_list: Optional[List[SequenceGroupMetadata]] = None
+ scheduler_outputs: Optional[SchedulerOutputs] = None
+
+ def as_broadcastable_tensor_dict(self) -> Dict[str, Any]:
+ tensor_dict = {
+ "input_tokens": self.input_tokens,
+ "input_positions": self.input_positions,
+ "lora_requests": self.lora_requests,
+ "lora_mapping": self.lora_mapping,
+ "multi_modal_kwargs": self.multi_modal_kwargs,
+ "prompt_adapter_mapping": self.prompt_adapter_mapping,
+ "prompt_adapter_requests": self.prompt_adapter_requests,
+ "virtual_engine": self.virtual_engine,
+ "request_ids_to_seq_ids": self.request_ids_to_seq_ids,
+ "finished_requests_ids": self.finished_requests_ids,
+ }
+ _add_attn_metadata_broadcastable_dict(tensor_dict, self.attn_metadata)
+ return tensor_dict
+
+ @classmethod
+ def from_broadcasted_tensor_dict(
+ cls: Type[TModelInputForGPU],
+ tensor_dict: Dict[str, Any],
+ attn_backend: Optional["AttentionBackend"] = None,
+ ) -> TModelInputForGPU:
+ if attn_backend is not None:
+ tensor_dict = _init_attn_metadata_from_tensor_dict(
+ attn_backend, tensor_dict)
+ return cls(**tensor_dict)
+
+
+@dataclass(frozen=True)
+class ModelInputForGPUWithSamplingMetadata(ModelInputForGPU):
+ """
+ Used by the ModelRunner.
+ """
+ sampling_metadata: Optional["SamplingMetadata"] = None
+ # Used for speculative decoding. We do not broadcast it because it is only
+ # used by the driver worker.
+ is_prompt: Optional[bool] = None
+
+ def as_broadcastable_tensor_dict(self) -> Dict[str, Any]:
+ tensor_dict = {
+ "input_tokens": self.input_tokens,
+ "input_positions": self.input_positions,
+ "lora_requests": self.lora_requests,
+ "lora_mapping": self.lora_mapping,
+ "multi_modal_kwargs": self.multi_modal_kwargs,
+ "prompt_adapter_mapping": self.prompt_adapter_mapping,
+ "prompt_adapter_requests": self.prompt_adapter_requests,
+ "virtual_engine": self.virtual_engine,
+ "request_ids_to_seq_ids": self.request_ids_to_seq_ids,
+ "finished_requests_ids": self.finished_requests_ids,
+ }
+ _add_attn_metadata_broadcastable_dict(tensor_dict, self.attn_metadata)
+ _add_sampling_metadata_broadcastable_dict(tensor_dict,
+ self.sampling_metadata)
+ return tensor_dict
+
+ @classmethod
+ def from_broadcasted_tensor_dict(
+ cls,
+ tensor_dict: Dict[str, Any],
+ attn_backend: Optional["AttentionBackend"] = None,
+ ) -> "ModelInputForGPUWithSamplingMetadata":
+ tensor_dict = _init_sampling_metadata_from_tensor_dict(tensor_dict)
+ if attn_backend is not None:
+ tensor_dict = _init_attn_metadata_from_tensor_dict(
+ attn_backend, tensor_dict)
+ return cls(**tensor_dict)
+
+
+class ModelInputForGPUBuilder(ModelRunnerInputBuilderBase[ModelInputForGPU]):
+ """Build ModelInputForGPU from SequenceGroupMetadata."""
+
+ # Note: ideally we would be using a dataclass(kw_only=True)
+ # here, so that this can be subclassed easily,
+ # but kw_only is not supported in python<3.10.
+ class InterDataForSeqGroup:
+ """Intermediate data for the current sequence group."""
+
+ def simple_reinit(self):
+ self.input_tokens[0].clear() # type: ignore
+ self.input_positions[0].clear() # type: ignore
+ self.mrope_input_positions = None # type: ignore
+ self.seq_lens[0] = 0 # type: ignore
+ self.orig_seq_lens[0] = 0 # type: ignore
+ self.query_lens[0] = 0 # type: ignore
+ self.context_lens[0] = 0 # type: ignore
+ self.curr_sliding_window_blocks[0] = 0 # type: ignore
+ self.lora_index_mapping.clear() # type: ignore
+ self.lora_prompt_mapping.clear() # type: ignore
+ self.lora_requests.clear() # type: ignore
+ self.prompt_adapter_index_mapping.clear() # type: ignore
+ self.prompt_adapter_prompt_mapping.clear() # type: ignore
+
+ def __init__(
+ self,
+ *,
+ # From sequence group metadata.
+ request_id: str,
+ seq_ids: List[int],
+ is_prompt: bool,
+ block_tables: Optional[Dict[int, List[int]]],
+ computed_block_nums: List[int],
+ n_seqs: int = 0,
+
+ # Input tokens and positions.
+ input_tokens: Optional[List[List[int]]] = None,
+ input_positions: Optional[List[List[int]]] = None,
+ mrope_input_positions: Optional[List[List[List[int]]]] = None,
+
+ # The sequence length (may be capped to the sliding window).
+ seq_lens: Optional[List[int]] = None,
+ # The original sequence length (before applying sliding window).
+ # This is used to compute slot mapping.
+ orig_seq_lens: Optional[List[int]] = None,
+ # The query length.
+ query_lens: Optional[List[int]] = None,
+ # The number of tokens that are already computed.
+ context_lens: Optional[List[int]] = None,
+ # The current sliding window block.
+ curr_sliding_window_blocks: Optional[List[int]] = None,
+
+ # LoRA inputs.
+ lora_index_mapping: Optional[List[List[int]]] = None,
+ lora_prompt_mapping: Optional[List[List[int]]] = None,
+ lora_requests: Optional[Set[LoRARequest]] = None,
+
+ # Prompt adapter inputs.
+ prompt_adapter_index_mapping: Optional[List[int]] = None,
+ prompt_adapter_prompt_mapping: Optional[List[int]] = None,
+ prompt_adapter_request: Optional[PromptAdapterRequest] = None,
+
+ # Multi-modal inputs.
+ multi_modal_inputs: Optional[MultiModalInputs] = None,
+
+ # Whether the prefix cache is hit (prefill only).
+ prefix_cache_hit: bool = False,
+ reinit: bool = False,
+ reinit_use_defaults: bool = False,
+ encoder_seq_len: int = 0,
+ ):
+ if reinit:
+ assert len(self.seq_ids) == len(seq_ids) # type: ignore
+ for i, seq_id in enumerate(seq_ids):
+ self.seq_ids[i] = seq_id # type: ignore
+ else:
+ self.seq_ids = seq_ids
+
+ self.request_id = request_id
+ self.is_prompt = is_prompt
+ self.block_tables = block_tables
+ self.computed_block_nums = computed_block_nums
+ self.n_seqs = n_seqs
+ self.encoder_seq_len = encoder_seq_len
+
+ if reinit:
+ if len(self.seq_ids) == 1 and reinit_use_defaults:
+ self.simple_reinit()
+ else:
+ if input_tokens:
+ self.input_tokens = input_tokens
+ else:
+ for seq_id in range(len(self.seq_ids)):
+ self.input_tokens[seq_id].clear()
+
+ if input_positions:
+ self.input_positions = input_positions
+ else:
+ for seq_id in range(len(self.seq_ids)):
+ self.input_positions[seq_id].clear()
+
+ self.mrope_input_positions = None
+
+ if seq_lens:
+ self.seq_lens = seq_lens
+ else:
+ for seq_id in range(len(self.seq_ids)):
+ self.seq_lens[seq_id] = 0
+
+ if orig_seq_lens:
+ self.orig_seq_lens = orig_seq_lens
+ else:
+ for seq_id in range(len(self.seq_ids)):
+ self.orig_seq_lens[seq_id] = 0
+
+ if query_lens:
+ self.query_lens = query_lens
+ else:
+ for seq_id in range(len(self.seq_ids)):
+ self.query_lens[seq_id] = 0
+
+ if context_lens:
+ self.context_lens = context_lens
+ else:
+ for seq_id in range(len(self.seq_ids)):
+ self.context_lens[seq_id] = 0
+
+ if curr_sliding_window_blocks:
+ self.curr_sliding_window_blocks = \
+ curr_sliding_window_blocks
+ else:
+ for seq_id in range(len(self.seq_ids)):
+ self.curr_sliding_window_blocks[seq_id] = 0
+
+ if lora_index_mapping:
+ self.lora_index_mapping = lora_index_mapping
+ else:
+ self.lora_index_mapping.clear()
+
+ if lora_prompt_mapping:
+ self.lora_prompt_mapping = lora_prompt_mapping
+ else:
+ self.lora_prompt_mapping.clear()
+
+ if lora_requests:
+ self.lora_requests = lora_requests
+ else:
+ self.lora_requests.clear()
+
+ if prompt_adapter_index_mapping:
+ self.prompt_adapter_index_mapping = \
+ prompt_adapter_index_mapping
+ else:
+ self.prompt_adapter_index_mapping.clear()
+
+ if prompt_adapter_prompt_mapping:
+ self.prompt_adapter_prompt_mapping = \
+ prompt_adapter_prompt_mapping
+ else:
+ self.prompt_adapter_prompt_mapping.clear()
+
+ else:
+ self.input_tokens = input_tokens or []
+ self.input_positions = input_positions or []
+ self.mrope_input_positions = mrope_input_positions or None
+ self.seq_lens = seq_lens or []
+ self.orig_seq_lens = orig_seq_lens or []
+ self.query_lens = query_lens or []
+ self.context_lens = context_lens or []
+ self.curr_sliding_window_blocks = \
+ curr_sliding_window_blocks or []
+
+ self.lora_index_mapping = lora_index_mapping or []
+ self.lora_prompt_mapping = lora_prompt_mapping or []
+ self.lora_requests = lora_requests or set()
+
+ self.prompt_adapter_index_mapping = (
+ prompt_adapter_index_mapping or [])
+ self.prompt_adapter_prompt_mapping = (
+ prompt_adapter_prompt_mapping or [])
+
+ self.prompt_adapter_request = prompt_adapter_request
+ self.multi_modal_inputs = multi_modal_inputs
+ self.prefix_cache_hit = prefix_cache_hit
+
+ self.n_seqs = len(self.seq_ids)
+
+ if not reinit:
+ self.__post_init__()
+
+ def __post_init__(self):
+ self.n_seqs = len(self.seq_ids)
+
+ self.input_tokens = [[] for _ in range(self.n_seqs)]
+ self.input_positions = [[] for _ in range(self.n_seqs)]
+ self.mrope_input_positions = None
+ self.seq_lens = [0] * self.n_seqs
+ self.orig_seq_lens = [0] * self.n_seqs
+ self.query_lens = [0] * self.n_seqs
+ self.context_lens = [0] * self.n_seqs
+ self.curr_sliding_window_blocks = [0] * self.n_seqs
+
+ self.lora_index_mapping = []
+ self.lora_prompt_mapping = []
+
+ def gen_inter_data_builder(self, num_seqs: int):
+ return lambda: ModelInputForGPUBuilder.InterDataForSeqGroup(
+ request_id="",
+ seq_ids=[0] * num_seqs,
+ is_prompt=True,
+ block_tables=None,
+ computed_block_nums=[])
+
+ def init_cached_inter_data(self, *args, **kwargs):
+ assert len(args) == 0
+ assert "seq_ids" in kwargs
+ seq_ids = kwargs["seq_ids"]
+ num_seqs = len(seq_ids)
+
+ # The inter-data cache is per model_runner
+ inter_data_cache = self.runner.inter_data_cache
+ if num_seqs not in inter_data_cache:
+ inter_data_cache[num_seqs] = PyObjectCache(
+ self.gen_inter_data_builder(num_seqs))
+
+ obj = inter_data_cache[num_seqs].get_object()
+ obj.__init__(*args, **kwargs)
+ return obj
+
+ def reset_cached_inter_data(self):
+ for cache in self.runner.inter_data_cache.values():
+ cache.reset()
+
+ def __init__(self,
+ runner: "GPUModelRunnerBase",
+ finished_requests_ids: Optional[List[str]] = None):
+ super().__init__()
+ # Compute functions for each sequence in a sequence group.
+ # WARNING: The order of the functions matters!
+ self.per_seq_compute_fns = [
+ self._compute_lens,
+ self._compute_for_prefix_cache_hit,
+ self._compute_for_sliding_window,
+ self._compute_lora_input,
+ ]
+ # Compute functions for each sequence group.
+ # WARNING: The order of the functions matters!
+ self.per_seq_group_compute_fns = [
+ self._compute_prompt_adapter_input,
+ self._compute_multi_modal_input,
+ ]
+
+ self.runner = runner
+ self.model_input_cls = self.runner._model_input_cls
+ self.attn_backend = self.runner.attn_backend
+ self.scheduler_config = self.runner.scheduler_config
+ self.sliding_window = self.runner.sliding_window
+ self.block_size = self.runner.block_size
+ self.enable_lora = self.runner.lora_config is not None
+ self.enable_prompt_adapter = (self.runner.prompt_adapter_config
+ is not None)
+ self.multi_modal_input_mapper = self.runner.multi_modal_input_mapper
+ self.finished_requests_ids = finished_requests_ids
+ self.decode_only = True
+
+ # Intermediate data (data in CPU before going to GPU) for
+ # the current sequence group.
+ self.inter_data_list: List[
+ ModelInputForGPUBuilder.InterDataForSeqGroup] = []
+
+ # Attention metadata inputs.
+ self.attn_metadata_builder = self.attn_backend.make_metadata_builder(
+ weakref.proxy(self))
+
+ # Engine/Model configurations.
+ self.chunked_prefill_enabled = (
+ self.scheduler_config is not None
+ and self.scheduler_config.chunked_prefill_enabled)
+ if self.sliding_window is not None:
+ self.sliding_window_blocks = (
+ self.sliding_window + self.block_size - 1) // self.block_size
+ self.block_aligned_sliding_window = \
+ self.sliding_window_blocks * self.block_size
+
+ def _compute_lens(self, inter_data: InterDataForSeqGroup, seq_idx: int,
+ seq_group_metadata: SequenceGroupMetadata):
+ """Compute context length, sequence length and tokens
+ for the given sequence data.
+ """
+ seq_data = seq_group_metadata.seq_data[inter_data.seq_ids[seq_idx]]
+ token_chunk_size = seq_group_metadata.token_chunk_size
+
+ # Compute context length (the number of tokens that are
+ # already computed) and sequence length (total number of tokens).
+
+ seq_len = seq_data.get_len()
+ if inter_data.is_prompt:
+ context_len = seq_data.get_num_computed_tokens()
+ seq_len = min(seq_len, context_len + token_chunk_size)
+ elif self.runner.scheduler_config.is_multi_step or \
+ self.runner.model_config.is_encoder_decoder_model:
+ context_len = seq_len - 1
+ else:
+ context_len = seq_data.get_num_computed_tokens()
+
+ # Compute tokens.
+ tokens = seq_data.get_token_ids()[context_len:seq_len]
+
+ inter_data.seq_lens[seq_idx] = seq_len
+ inter_data.orig_seq_lens[seq_idx] = seq_len
+ inter_data.context_lens[seq_idx] = context_len
+ inter_data.input_tokens[seq_idx].extend(tokens)
+ inter_data.input_positions[seq_idx].extend(range(context_len, seq_len))
+ inter_data.query_lens[seq_idx] = seq_len - context_len
+
+ if seq_data.mrope_position_delta is not None:
+ if inter_data.mrope_input_positions is None:
+ inter_data.mrope_input_positions = [None] * inter_data.n_seqs
+
+ inter_data.mrope_input_positions[
+ seq_idx] = MRotaryEmbedding.get_next_input_positions(
+ seq_data.mrope_position_delta,
+ context_len,
+ seq_len,
+ )
+
+ def _compute_for_prefix_cache_hit(
+ self, inter_data: InterDataForSeqGroup, seq_idx: int,
+ seq_group_metadata: SequenceGroupMetadata):
+ """Check if hit prefix cache (i.e., some blocks are already computed).
+ If hit, update input tokens and positions to only compute the
+ remaining blocks.
+ """
+ computed_block_nums = inter_data.computed_block_nums
+
+ # Note that prefix caching does not support sliding window.
+ prefix_cache_hit = (computed_block_nums is not None
+ and len(computed_block_nums) > 0
+ and self.sliding_window is None
+ and inter_data.is_prompt)
+ inter_data.prefix_cache_hit = prefix_cache_hit
+
+ if not prefix_cache_hit:
+ return
+
+ assert computed_block_nums is not None
+ # The cache hit prompt tokens in this sequence. Note that
+ # this may be larger than the sequence length if chunked
+ # prefill is enabled.
+ prefix_cache_len = len(computed_block_nums) * self.block_size
+ # The number of so far computed prompt tokens in this sequence.
+ context_len = inter_data.context_lens[seq_idx]
+ # The total number of prompt tokens in this sequence.
+ # When chunked prefill is enabled, this is the token number of
+ # computed chunks + current chunk.
+ seq_len = inter_data.seq_lens[seq_idx]
+ if prefix_cache_len <= context_len:
+ # We already passed the cache hit region,
+ # so do normal computation.
+ # Must clear prefix_cache_hit so _add_seq_group uses the full
+ # block_tables (prefix + previous-chunk blocks) instead of only
+ # computed_block_nums (prefix only). Without this, block_tables
+ # passed to _forward_prefix_pytorch is too narrow for context_len,
+ # causing an empty blk_ids slice and a zero-dim amax() crash.
+ inter_data.prefix_cache_hit = False
+ elif context_len < prefix_cache_len < seq_len:
+ # Partial hit. Compute the missing part.
+ uncomputed_start = prefix_cache_len - context_len
+ inter_data.input_tokens[seq_idx] = inter_data.input_tokens[
+ seq_idx][uncomputed_start:]
+ inter_data.input_positions[seq_idx] = inter_data.input_positions[
+ seq_idx][uncomputed_start:]
+ context_len = prefix_cache_len
+
+ inter_data.context_lens[seq_idx] = context_len
+ inter_data.query_lens[
+ seq_idx] = inter_data.seq_lens[seq_idx] - context_len
+ elif seq_len <= prefix_cache_len:
+ # Full hit. Only compute the last token to avoid
+ # erroneous behavior. FIXME: Ideally we should directly
+ # mark all tokens as computed in the scheduler and do not
+ # schedule this sequence, so this case should not happen.
+ inter_data.input_tokens[seq_idx] = inter_data.input_tokens[
+ seq_idx][-1:]
+ inter_data.input_positions[seq_idx] = inter_data.input_positions[
+ seq_idx][-1:]
+ inter_data.query_lens[seq_idx] = 1
+ inter_data.context_lens[seq_idx] = inter_data.seq_lens[seq_idx] - 1
+
+ def _compute_for_sliding_window(self, inter_data: InterDataForSeqGroup,
+ seq_idx: int,
+ seq_group_metadata: SequenceGroupMetadata):
+ """Update seq_len and curr_sliding_window_block for the given
+ sequence data (only required by decoding) if sliding window is enabled.
+ """
+ curr_sliding_window_block = 0
+ sliding_seq_len = inter_data.seq_lens[seq_idx]
+ if not inter_data.is_prompt and self.sliding_window is not None:
+ # TODO(sang): This is a hack to make sliding window work with
+ # paged attn. We can remove it if we make paged attn kernel
+ # to properly handle slinding window attn.
+ curr_sliding_window_block = self.sliding_window_blocks
+ if self.scheduler_config.use_v2_block_manager:
+ # number of elements in last block
+ suff_len = inter_data.seq_lens[seq_idx] % self.block_size
+ sliding_seq_len = min(
+ inter_data.seq_lens[seq_idx],
+ self.block_aligned_sliding_window + suff_len)
+ if suff_len > 0:
+ curr_sliding_window_block += 1
+ else:
+ sliding_seq_len = min(inter_data.seq_lens[seq_idx],
+ self.sliding_window)
+
+ inter_data.curr_sliding_window_blocks[
+ seq_idx] = curr_sliding_window_block
+ inter_data.seq_lens[seq_idx] = sliding_seq_len
+
+ def _compute_lora_input(self, inter_data: InterDataForSeqGroup,
+ seq_idx: int,
+ seq_group_metadata: SequenceGroupMetadata):
+ """If LoRA is enabled, compute LoRA index and prompt mapping."""
+ if not self.enable_lora:
+ return
+
+ lora_id = seq_group_metadata.lora_int_id
+ if lora_id > 0:
+ inter_data.lora_requests.add(seq_group_metadata.lora_request)
+ query_len = inter_data.query_lens[seq_idx]
+ inter_data.lora_index_mapping.append([lora_id] * query_len)
+ inter_data.lora_prompt_mapping.append(
+ [lora_id] *
+ (query_len if seq_group_metadata.sampling_params
+ and seq_group_metadata.sampling_params.prompt_logprobs is not None
+ else 1))
+
+ def _compute_prompt_adapter_input(
+ self, inter_data: InterDataForSeqGroup,
+ seq_group_metadata: SequenceGroupMetadata):
+ """If prompt adapter is enabled, compute index and prompt mapping.
+ """
+ # Note that when is_prompt=True, we expect only one sequence
+ # in the group.
+ if not self.enable_prompt_adapter:
+ return
+
+ prompt_adapter_id = seq_group_metadata.prompt_adapter_id
+ if prompt_adapter_id <= 0 or not inter_data.is_prompt:
+ return
+
+ # We expect only one sequence in the group when is_prompt=True.
+ assert inter_data.n_seqs == 1
+ query_len = inter_data.query_lens[0]
+ inter_data.prompt_adapter_request = (
+ seq_group_metadata.prompt_adapter_request)
+
+ num_tokens = seq_group_metadata.prompt_adapter_num_virtual_tokens
+ inter_data.prompt_adapter_index_mapping = [
+ prompt_adapter_id
+ ] * num_tokens + [0] * (query_len - num_tokens)
+ inter_data.prompt_adapter_prompt_mapping = [prompt_adapter_id] * (
+ query_len if seq_group_metadata.sampling_params
+ and seq_group_metadata.sampling_params.prompt_logprobs else 1)
+
+ def _compute_multi_modal_input(self, inter_data: InterDataForSeqGroup,
+ seq_group_metadata: SequenceGroupMetadata):
+ """If multi-modal data is given, add it to the input."""
+ mm_data = seq_group_metadata.multi_modal_data
+ if not mm_data:
+ return
+
+ mm_kwargs = self.multi_modal_input_mapper(
+ mm_data,
+ mm_processor_kwargs=seq_group_metadata.mm_processor_kwargs)
+ inter_data.multi_modal_inputs = mm_kwargs
+
+ # special processing for mrope position deltas.
+ if self.runner.model_is_mrope:
+ image_grid_thw = mm_kwargs.get("image_grid_thw", None)
+ video_grid_thw = mm_kwargs.get("video_grid_thw", None)
+ assert image_grid_thw is not None or video_grid_thw is not None, (
+ "mrope embedding type requires multi-modal input mapper "
+ "returns 'image_grid_thw' or 'video_grid_thw'.")
+
+ hf_config = self.runner.model_config.hf_config
+
+ inter_data.mrope_input_positions = [None] * inter_data.n_seqs
+ for seq_idx in range(inter_data.n_seqs):
+ seq_data = seq_group_metadata.seq_data[
+ inter_data.seq_ids[seq_idx]]
+ token_ids = seq_data.get_token_ids()
+
+ mrope_input_positions, mrope_position_delta = \
+ MRotaryEmbedding.get_input_positions(
+ token_ids,
+ image_grid_thw=image_grid_thw,
+ video_grid_thw=video_grid_thw,
+ image_token_id=hf_config.image_token_id,
+ video_token_id=hf_config.video_token_id,
+ vision_start_token_id=hf_config.vision_start_token_id,
+ vision_end_token_id=hf_config.vision_end_token_id,
+ spatial_merge_size=hf_config.vision_config.
+ spatial_merge_size,
+ context_len=inter_data.context_lens[seq_idx],
+ )
+
+ seq_data.mrope_position_delta = mrope_position_delta
+ inter_data.mrope_input_positions[
+ seq_idx] = mrope_input_positions
+
+ def add_seq_group(self, seq_group_metadata: SequenceGroupMetadata):
+ """Add a sequence group to the builder."""
+ seq_ids = seq_group_metadata.seq_data.keys()
+ n_seqs = len(seq_ids)
+ is_prompt = seq_group_metadata.is_prompt
+
+ if is_prompt:
+ assert n_seqs == 1
+ self.decode_only = False
+
+ encoder_seq_len = 0
+
+ if self.runner.model_config.is_encoder_decoder_model:
+ encoder_seq_len = seq_group_metadata.encoder_seq_data.get_len()
+
+ inter_data = self.init_cached_inter_data(
+ request_id=seq_group_metadata.request_id,
+ seq_ids=seq_ids,
+ is_prompt=is_prompt,
+ block_tables=seq_group_metadata.block_tables,
+ computed_block_nums=seq_group_metadata.computed_block_nums,
+ reinit=True,
+ reinit_use_defaults=True,
+ encoder_seq_len=encoder_seq_len)
+
+ self.inter_data_list.append(inter_data)
+
+ for seq_idx in range(n_seqs):
+ for per_seq_fn in self.per_seq_compute_fns:
+ per_seq_fn(inter_data, seq_idx, seq_group_metadata)
+ for per_seq_group_fn in self.per_seq_group_compute_fns:
+ per_seq_group_fn(inter_data, seq_group_metadata)
+
+ def _use_captured_graph(self,
+ batch_size: int,
+ decode_only: bool,
+ max_decode_seq_len: int,
+ max_encoder_seq_len: int = 0) -> bool:
+ return (decode_only and not self.runner.model_config.enforce_eager
+ and batch_size <= _BATCH_SIZES_TO_CAPTURE[-1]
+ and max_decode_seq_len <= self.runner.max_seq_len_to_capture
+ and max_encoder_seq_len <= self.runner.max_seq_len_to_capture
+ and batch_size <= self.runner.max_batchsize_to_capture)
+
+ def _get_cuda_graph_pad_size(self,
+ num_seqs: int,
+ max_decode_seq_len: int,
+ max_encoder_seq_len: int = 0) -> int:
+ """
+ Determine the number of padding sequences required for running in
+ CUDA graph mode. Returns -1 if CUDA graphs cannot be used.
+
+ In the multi-step + chunked-prefill case, only the first step
+ has Prefills (if any). The rest of the steps are guaranteed to be all
+ decodes. In this case, we set up the padding as if all the sequences
+ are decodes so we may run all steps except the first step in CUDA graph
+ mode. The padding is accounted for in the multi-step `advance_step`
+ family of functions.
+
+ Args:
+ num_seqs (int): Number of sequences scheduled to run.
+ max_decode_seq_len (int): Greatest of all the decode sequence
+ lengths. Used only in checking the viablility of using
+ CUDA graphs.
+ max_encoder_seq_len (int, optional): Greatest of all the encode
+ sequence lengths. Defaults to 0. Used only in checking the
+ viability of using CUDA graphs.
+ Returns:
+ int: Returns the determined number of padding sequences. If
+ CUDA graphs is not viable, returns -1.
+ """
+ is_mscp: bool = self.runner.scheduler_config.is_multi_step and \
+ self.runner.scheduler_config.chunked_prefill_enabled
+ decode_only = self.decode_only or is_mscp
+ if not decode_only:
+ # Early exit so we can treat num_seqs as the batch_size below.
+ return -1
+
+ # batch_size out of this function refers to the number of input
+ # tokens being scheduled. This conflation of num_seqs as batch_size
+ # is valid as this is a decode-only case.
+ batch_size = num_seqs
+ if not self._use_captured_graph(batch_size, decode_only,
+ max_decode_seq_len,
+ max_encoder_seq_len):
+ return -1
+
+ graph_batch_size = _get_graph_batch_size(batch_size)
+ assert graph_batch_size >= batch_size
+ return graph_batch_size - batch_size
+
+ def build(self) -> ModelInputForGPU:
+ """Finalize the builder intermediate data and
+ create on-device tensors.
+ """
+ # Combine and flatten intermediate data.
+ input_tokens = []
+ for inter_data in self.inter_data_list:
+ for cur_input_tokens in inter_data.input_tokens:
+ input_tokens.extend(cur_input_tokens)
+
+ if not input_tokens:
+ # This may happen when all prefill requests hit
+ # prefix caching and there is no decode request.
+ return self.model_input_cls()
+
+ mrope_input_positions: Optional[List[List[int]]] = None
+ if any(inter_data.mrope_input_positions is not None
+ for inter_data in self.inter_data_list):
+ mrope_input_positions = [[] for _ in range(3)]
+ for idx in range(3):
+ for inter_data in self.inter_data_list:
+ msections = inter_data.mrope_input_positions
+ if msections is None:
+ for _seq_input_positions in inter_data.input_positions:
+ mrope_input_positions[idx].extend(
+ _seq_input_positions)
+ else:
+ for _seq_mrope_input_positions in msections:
+ mrope_input_positions[idx].extend(
+ _seq_mrope_input_positions[idx])
+ input_positions = None
+ else:
+ input_positions = []
+ for inter_data in self.inter_data_list:
+ for cur_input_positions in inter_data.input_positions:
+ input_positions.extend(cur_input_positions)
+
+ seq_lens = []
+ query_lens = []
+ max_decode_seq_len = 0
+ max_encoder_seq_len = 0
+ for inter_data in self.inter_data_list:
+ seq_lens.extend(inter_data.seq_lens)
+ query_lens.extend(inter_data.query_lens)
+ if not inter_data.is_prompt:
+ max_decode_seq_len = max(max_decode_seq_len,
+ max(inter_data.seq_lens))
+ if self.runner.model_config.is_encoder_decoder_model:
+ max_encoder_seq_len = max(max_encoder_seq_len,
+ inter_data.encoder_seq_len)
+
+ # Mapping from request IDs to sequence IDs. Used for Jamba models
+ # that manages the cache by itself.
+ request_ids_to_seq_ids = {
+ data.request_id: data.seq_ids
+ for data in self.inter_data_list
+ }
+
+ cuda_graph_pad_size = self._get_cuda_graph_pad_size(
+ num_seqs=len(seq_lens),
+ max_decode_seq_len=max_encoder_seq_len,
+ max_encoder_seq_len=max_encoder_seq_len)
+
+ batch_size = len(input_tokens)
+ if cuda_graph_pad_size != -1:
+ # If cuda graph can be used, pad tensors accordingly.
+ # See `capture_model` API for more details.
+ # vLLM uses cuda graph only for decoding requests.
+ batch_size += cuda_graph_pad_size
+
+ # Tokens and positions.
+ if cuda_graph_pad_size:
+ input_tokens.extend(itertools.repeat(0, cuda_graph_pad_size))
+ assert self.runner.device is not None
+ input_tokens_tensor = async_tensor_h2d(input_tokens, torch.long,
+ self.runner.device,
+ self.runner.pin_memory)
+ if mrope_input_positions is not None:
+ for idx in range(3):
+ mrope_input_positions[idx].extend(
+ itertools.repeat(0, cuda_graph_pad_size))
+ input_positions_tensor = async_tensor_h2d(mrope_input_positions,
+ torch.long,
+ self.runner.device,
+ self.runner.pin_memory)
+ else:
+ input_positions.extend(itertools.repeat(0, cuda_graph_pad_size))
+ input_positions_tensor = async_tensor_h2d(input_positions,
+ torch.long,
+ self.runner.device,
+ self.runner.pin_memory)
+ # Sequence and query lengths.
+ if cuda_graph_pad_size:
+ seq_lens.extend(itertools.repeat(1, cuda_graph_pad_size))
+
+ # Attention metadata.
+ attn_metadata = self.attn_metadata_builder.build(
+ seq_lens, query_lens, cuda_graph_pad_size, batch_size)
+
+ # LoRA data.
+ lora_requests = set()
+ lora_mapping = None
+ if self.enable_lora:
+ lora_requests = set(r for data in self.inter_data_list
+ for r in data.lora_requests)
+ lora_index_mapping = flatten_2d_lists([
+ flatten_2d_lists(inter_data.lora_index_mapping)
+ for inter_data in self.inter_data_list
+ ])
+ if cuda_graph_pad_size:
+ lora_index_mapping.extend(
+ itertools.repeat(0, cuda_graph_pad_size))
+ lora_prompt_mapping = flatten_2d_lists([
+ flatten_2d_lists(inter_data.lora_prompt_mapping)
+ for inter_data in self.inter_data_list
+ ])
+
+ lora_mapping = LoRAMapping(
+ **dict(index_mapping=lora_index_mapping,
+ prompt_mapping=lora_prompt_mapping,
+ is_prefill=not self.decode_only))
+
+ # Prompt adapter data.
+ prompt_adapter_requests: Set[PromptAdapterRequest] = set()
+ prompt_adapter_mapping = None
+ if self.enable_prompt_adapter:
+ prompt_adapter_requests = set(
+ data.prompt_adapter_request for data in self.inter_data_list
+ if data.prompt_adapter_request is not None)
+ prompt_adapter_index_mapping = flatten_2d_lists([
+ inter_data.prompt_adapter_index_mapping
+ for inter_data in self.inter_data_list
+ ])
+ if cuda_graph_pad_size:
+ prompt_adapter_index_mapping.extend(
+ itertools.repeat(0, cuda_graph_pad_size))
+ prompt_adapter_prompt_mapping = flatten_2d_lists([
+ inter_data.prompt_adapter_prompt_mapping
+ for inter_data in self.inter_data_list
+ ])
+ prompt_adapter_mapping = PromptAdapterMapping(
+ prompt_adapter_index_mapping,
+ prompt_adapter_prompt_mapping,
+ )
+
+ # Multi-modal data.
+ multi_modal_inputs_list = [
+ data.multi_modal_inputs for data in self.inter_data_list
+ if data.multi_modal_inputs is not None
+ ]
+ multi_modal_kwargs = MultiModalInputs.batch(multi_modal_inputs_list)
+
+ return self.model_input_cls(
+ input_tokens=input_tokens_tensor,
+ input_positions=input_positions_tensor,
+ attn_metadata=attn_metadata,
+ seq_lens=seq_lens,
+ query_lens=query_lens,
+ lora_mapping=lora_mapping,
+ lora_requests=lora_requests,
+ multi_modal_kwargs=multi_modal_kwargs,
+ request_ids_to_seq_ids=request_ids_to_seq_ids,
+ finished_requests_ids=self.finished_requests_ids,
+ prompt_adapter_mapping=prompt_adapter_mapping,
+ prompt_adapter_requests=prompt_adapter_requests)
+
+
+class GPUModelRunnerBase(ModelRunnerBase[TModelInputForGPU]):
+ """
+ Helper class for shared methods between GPU model runners.
+ """
+ _model_input_cls: Type[TModelInputForGPU]
+ _builder_cls: Type[ModelInputForGPUBuilder]
+
+ def __init__(
+ self,
+ model_config: ModelConfig,
+ parallel_config: ParallelConfig,
+ scheduler_config: SchedulerConfig,
+ device_config: DeviceConfig,
+ cache_config: CacheConfig,
+ load_config: LoadConfig,
+ lora_config: Optional[LoRAConfig],
+ kv_cache_dtype: Optional[str] = "auto",
+ is_driver_worker: bool = False,
+ prompt_adapter_config: Optional[PromptAdapterConfig] = None,
+ return_hidden_states: bool = False,
+ observability_config: Optional[ObservabilityConfig] = None,
+ input_registry: InputRegistry = INPUT_REGISTRY,
+ mm_registry: MultiModalRegistry = MULTIMODAL_REGISTRY,
+ ):
+ self.model_config = model_config
+ self.parallel_config = parallel_config
+ self.scheduler_config = scheduler_config
+ self.device_config = device_config
+ self.cache_config = cache_config
+ self.lora_config = lora_config
+ self.load_config = load_config
+ self.is_driver_worker = is_driver_worker
+ self.prompt_adapter_config = prompt_adapter_config
+ self.return_hidden_states = return_hidden_states
+ self.observability_config = observability_config
+
+ self.device = self.device_config.device
+ self.pin_memory = is_pin_memory_available()
+
+ self.kv_cache_dtype = kv_cache_dtype
+ self.sliding_window = model_config.get_sliding_window()
+ self.block_size = cache_config.block_size
+ self.max_seq_len_to_capture = self.model_config.max_seq_len_to_capture
+ self.max_batchsize_to_capture = _get_max_graph_batch_size(
+ self.scheduler_config.max_num_seqs)
+
+ self.graph_runners: List[Dict[int, CUDAGraphRunner]] = [
+ {} for _ in range(self.parallel_config.pipeline_parallel_size)
+ ]
+ self.graph_memory_pool: Optional[Tuple[
+ int, int]] = None # Set during graph capture.
+
+ self.has_inner_state = model_config.has_inner_state
+
+ # When using CUDA graph, the input block tables must be padded to
+ # max_seq_len_to_capture. However, creating the block table in
+ # Python can be expensive. To optimize this, we cache the block table
+ # in numpy and only copy the actual input content at every iteration.
+ # The shape of the cached block table will be
+ # (max batch size to capture, max context len to capture / block size).
+ self.graph_block_tables = np.zeros(
+ (self.max_batchsize_to_capture, self.get_max_block_per_batch()),
+ dtype=np.int32)
+
+ # Attention-free but stateful models like Mamba need a placeholder attn
+ # backend, as the attention metadata is needed to manage internal state.
+ # However we must bypass attention selection altogether for some models
+ # used for speculative decoding to avoid a divide-by-zero in
+ # model_config.get_head_size()
+ num_attn_heads = self.model_config.get_num_attention_heads(
+ self.parallel_config)
+ needs_attn_backend = (num_attn_heads != 0
+ or self.model_config.is_attention_free)
+
+ self.attn_backend = get_attn_backend(
+ self.model_config.get_head_size(),
+ self.model_config.get_sliding_window(),
+ self.model_config.dtype,
+ self.kv_cache_dtype,
+ self.block_size,
+ self.model_config.is_attention_free,
+ ) if needs_attn_backend else None
+ if self.attn_backend:
+ self.attn_state = self.attn_backend.get_state_cls()(
+ weakref.proxy(self))
+ else:
+ self.attn_state = CommonAttentionState(weakref.proxy(self))
+
+ # Multi-modal data support
+ self.input_registry = input_registry
+ self.mm_registry = mm_registry
+ self.multi_modal_input_mapper = mm_registry \
+ .create_input_mapper(model_config)
+ self.mm_registry.init_mm_limits_per_prompt(self.model_config)
+
+ # Lazy initialization
+ self.model: nn.Module # Set after load_model
+ # Set after load_model.
+ self.lora_manager: Optional[LRUCacheWorkerLoRAManager] = None
+ self.prompt_adapter_manager: LRUCacheWorkerPromptAdapterManager = None
+
+ set_cpu_offload_max_bytes(
+ int(self.cache_config.cpu_offload_gb * 1024**3))
+
+ # Used to cache python objects
+ self.inter_data_cache: Dict[int, PyObjectCache] = {}
+
+ # Using the PythonizationCache in Pipeline-Parallel clobbers the
+ # SequenceGroupToSample object. In Pipeline-Parallel, we have
+ # more than 1 Scheduler, resulting in a potential back-to-back
+ # prepare_model_inputs() call. This clobbers the cached
+ # SequenceGroupToSample objects, as we reset the cache during
+ # every prepare_model_inputs() call.
+ self.sampling_metadata_cache: SamplingMetadataCache = \
+ SamplingMetadataCache() \
+ if self.parallel_config.pipeline_parallel_size == 1 else None
+
+ def load_model(self) -> None:
+ logger.info("Starting to load model %s...", self.model_config.model)
+ with DeviceMemoryProfiler() as m:
+ self.model = get_model(model_config=self.model_config,
+ device_config=self.device_config,
+ load_config=self.load_config,
+ lora_config=self.lora_config,
+ parallel_config=self.parallel_config,
+ scheduler_config=self.scheduler_config,
+ cache_config=self.cache_config)
+
+ self.model_memory_usage = m.consumed_memory
+ logger.info("Loading model weights took %.4f GB",
+ self.model_memory_usage / float(2**30))
+
+ if self.lora_config:
+ assert supports_lora(
+ self.model
+ ), f"{self.model.__class__.__name__} does not support LoRA yet."
+
+ if supports_multimodal(self.model):
+ logger.warning("Regarding multimodal models, vLLM currently "
+ "only supports adding LoRA to language model.")
+ # It's necessary to distinguish between the max_position_embeddings
+ # of VLMs and LLMs.
+ if hasattr(self.model.config, "max_position_embeddings"):
+ max_pos_embeddings = self.model.config.max_position_embeddings
+ else:
+ max_pos_embeddings = (
+ self.model.config.text_config.max_position_embeddings)
+
+ self.lora_manager = LRUCacheWorkerLoRAManager(
+ self.scheduler_config.max_num_seqs,
+ self.scheduler_config.max_num_batched_tokens,
+ self.vocab_size,
+ self.lora_config,
+ self.device,
+ self.model.embedding_modules,
+ self.model.embedding_padding_modules,
+ max_position_embeddings=max_pos_embeddings,
+ )
+ self.model = self.lora_manager.create_lora_manager(self.model)
+
+ if self.prompt_adapter_config:
+ self.prompt_adapter_manager = LRUCacheWorkerPromptAdapterManager(
+ self.scheduler_config.max_num_seqs,
+ self.scheduler_config.max_num_batched_tokens, self.device,
+ self.prompt_adapter_config)
+ self.model = (
+ self.prompt_adapter_manager.create_prompt_adapter_manager(
+ self.model))
+
+ if self.kv_cache_dtype == "fp8" and is_hip():
+ # Currently only ROCm accepts kv-cache scaling factors
+ # via quantization_param_path and this will be deprecated
+ # in the future.
+ if self.model_config.quantization_param_path is not None:
+ if callable(getattr(self.model, "load_kv_cache_scales", None)):
+ warnings.warn(
+ "Loading kv cache scaling factor from JSON is "
+ "deprecated and will be removed. Please include "
+ "kv cache scaling factors in the model checkpoint.",
+ FutureWarning,
+ stacklevel=2)
+ self.model.load_kv_cache_scales(
+ self.model_config.quantization_param_path)
+ logger.info("Loaded KV cache scaling factors from %s",
+ self.model_config.quantization_param_path)
+ else:
+ raise RuntimeError(
+ "Using FP8 KV cache and scaling factors provided but "
+ "model %s does not support loading scaling factors.",
+ self.model.__class__)
+ else:
+ logger.warning(
+ "Using FP8 KV cache but no scaling factors "
+ "provided. Defaulting to scaling factors of 1.0. "
+ "This may lead to less accurate results!")
+
+ if envs.VLLM_TORCH_COMPILE_LEVEL == CompilationLevel.DYNAMO_AS_IS \
+ and supports_dynamo():
+ from vllm.plugins import get_torch_compile_backend
+ backend = get_torch_compile_backend() or "eager"
+ self.model = torch.compile(
+ self.model,
+ fullgraph=envs.VLLM_TEST_DYNAMO_FULLGRAPH_CAPTURE,
+ backend=backend)
+
+ def save_sharded_state(
+ self,
+ path: str,
+ pattern: Optional[str] = None,
+ max_size: Optional[int] = None,
+ ) -> None:
+ from vllm.model_executor.model_loader.loader import ShardedStateLoader
+ ShardedStateLoader.save_model(
+ self.model,
+ path,
+ pattern=pattern,
+ max_size=max_size,
+ )
+
+ def save_tensorized_model(
+ self,
+ tensorizer_config: TensorizerConfig,
+ ) -> None:
+ from vllm.model_executor.model_loader.loader import TensorizerLoader
+ TensorizerLoader.save_model(
+ self.model,
+ tensorizer_config=tensorizer_config,
+ )
+
+ def get_max_block_per_batch(self) -> int:
+ block_size = self.block_size
+ return (self.max_seq_len_to_capture + block_size - 1) // block_size
+
+ def _prepare_model_input_tensors(
+ self,
+ seq_group_metadata_list: List[SequenceGroupMetadata],
+ finished_requests_ids: Optional[List[str]] = None
+ ) -> TModelInputForGPU:
+ """Helper method to prepare the model input based on a given sequence
+ group. Prepares metadata needed for the base model forward pass but not
+ metadata for possible additional steps, e.g., sampling.
+
+ The API assumes seq_group_metadata_list is sorted by prefill -> decode.
+
+ The result tensors and data structure also batches input in prefill
+ -> decode order. For example,
+
+ - input_tokens[:num_prefill_tokens] contains prefill tokens.
+ - input_tokens[num_prefill_tokens:] contains decode tokens.
+
+ If cuda graph is required, this API automatically pads inputs.
+ """
+ builder = self._builder_cls(weakref.proxy(self), finished_requests_ids)
+ for seq_group_metadata in seq_group_metadata_list:
+ builder.add_seq_group(seq_group_metadata)
+
+ builder.reset_cached_inter_data()
+
+ return builder.build() # type: ignore
+
+ @torch.inference_mode()
+ def profile_run(self) -> None:
+ # Enable top-k sampling to reflect the accurate memory usage.
+ sampling_params = SamplingParams(top_p=0.99, top_k=self.vocab_size - 1)
+ max_num_batched_tokens = self.scheduler_config.max_num_batched_tokens
+ max_num_seqs = self.scheduler_config.max_num_seqs
+ # This represents the maximum number of different requests
+ # that will have unique loras, an therefore the max amount of memory
+ # consumption create dummy lora request copies from the lora request
+ # passed in, which contains a lora from the lora warmup path.
+ dummy_lora_requests: List[LoRARequest] = []
+ dummy_lora_requests_per_seq: List[LoRARequest] = []
+ if self.lora_config:
+ assert self.lora_manager is not None
+ with self.lora_manager.dummy_lora_cache():
+ for idx in range(self.lora_config.max_loras):
+ lora_id = idx + 1
+ dummy_lora_request = LoRARequest(
+ lora_name=f"warmup_{lora_id}",
+ lora_int_id=lora_id,
+ lora_path="/not/a/real/path",
+ )
+ self.lora_manager.add_dummy_lora(dummy_lora_request,
+ rank=LORA_WARMUP_RANK)
+ dummy_lora_requests.append(dummy_lora_request)
+ dummy_lora_requests_per_seq = [
+ dummy_lora_requests[idx % len(dummy_lora_requests)]
+ for idx in range(max_num_seqs)
+ ]
+
+ # Profile memory usage with max_num_sequences sequences and the total
+ # number of tokens equal to max_num_batched_tokens.
+ seqs: List[SequenceGroupMetadata] = []
+ # Additional GPU memory may be needed for multi-modal encoding, which
+ # needs to be accounted for when calculating the GPU blocks for
+ # vLLM blocker manager.
+ # To exercise the worst scenario for GPU memory consumption,
+ # the number of seqs (batch_size) is chosen to maximize the number
+ # of images processed.
+
+ max_mm_tokens = self.mm_registry.get_max_multimodal_tokens(
+ self.model_config)
+ if max_mm_tokens > 0:
+ max_num_seqs_orig = max_num_seqs
+ max_num_seqs = min(max_num_seqs,
+ max_num_batched_tokens // max_mm_tokens)
+ if max_num_seqs < 1:
+ expr = (f"min({max_num_seqs_orig}, "
+ f"{max_num_batched_tokens} // {max_mm_tokens})")
+ logger.warning(
+ "Computed max_num_seqs (%s) to be less than 1. "
+ "Setting it to the minimum value of 1.", expr)
+ max_num_seqs = 1
+
+ batch_size = 0
+ for group_id in range(max_num_seqs):
+ seq_len = (max_num_batched_tokens // max_num_seqs +
+ (group_id < max_num_batched_tokens % max_num_seqs))
+ batch_size += seq_len
+
+ seq_data, dummy_multi_modal_data = self.input_registry \
+ .dummy_data_for_profiling(self.model_config,
+ seq_len,
+ self.mm_registry)
+
+ seq = SequenceGroupMetadata(
+ request_id=str(group_id),
+ is_prompt=True,
+ seq_data={group_id: seq_data},
+ sampling_params=sampling_params,
+ block_tables=None,
+ lora_request=dummy_lora_requests_per_seq[group_id]
+ if dummy_lora_requests_per_seq else None,
+ multi_modal_data=dummy_multi_modal_data,
+ )
+ seqs.append(seq)
+
+ # Run the model with the dummy inputs.
+ num_layers = self.model_config.get_num_layers(self.parallel_config)
+ # use an empty tensor instead of `None`` to force Dynamo to pass
+ # it by reference, rather by specializing on the value ``None``.
+ # the `dtype` argument does not matter, and we use `float32` as
+ # a placeholder (it has wide hardware support).
+ # it is important to create tensors inside the loop, rather than
+ # multiplying the list, to avoid Dynamo from treating them as
+ # tensor aliasing.
+ kv_caches = [
+ torch.tensor([], dtype=torch.float32, device=self.device)
+ for _ in range(num_layers)
+ ]
+ finished_requests_ids = [seq.request_id for seq in seqs]
+ model_input = self.prepare_model_input(
+ seqs, finished_requests_ids=finished_requests_ids)
+ intermediate_tensors = None
+ if not get_pp_group().is_first_rank:
+ intermediate_tensors = self.model.make_empty_intermediate_tensors(
+ batch_size=batch_size,
+ dtype=self.model_config.dtype,
+ device=self.device)
+
+ graph_batch_size = self.max_batchsize_to_capture
+ batch_size_capture_list = [
+ bs for bs in _BATCH_SIZES_TO_CAPTURE if bs <= graph_batch_size
+ ]
+ if self.model_config.enforce_eager:
+ batch_size_capture_list = []
+ with set_compile_context(batch_size_capture_list):
+ self.execute_model(model_input, kv_caches, intermediate_tensors)
+ torch.cuda.synchronize()
+ return
+
+ def remove_all_loras(self):
+ if not self.lora_manager:
+ raise RuntimeError("LoRA is not enabled.")
+ self.lora_manager.remove_all_adapters()
+
+ def set_active_loras(self, lora_requests: Set[LoRARequest],
+ lora_mapping: LoRAMapping) -> None:
+ if not self.lora_manager:
+ raise RuntimeError("LoRA is not enabled.")
+ self.lora_manager.set_active_adapters(lora_requests, lora_mapping)
+
+ def add_lora(self, lora_request: LoRARequest) -> bool:
+ if not self.lora_manager:
+ raise RuntimeError("LoRA is not enabled.")
+ return self.lora_manager.add_adapter(lora_request)
+
+ def remove_lora(self, lora_id: int) -> bool:
+ if not self.lora_manager:
+ raise RuntimeError("LoRA is not enabled.")
+ return self.lora_manager.remove_adapter(lora_id)
+
+ def pin_lora(self, lora_id: int) -> bool:
+ if not self.lora_manager:
+ raise RuntimeError("LoRA is not enabled.")
+ return self.lora_manager.pin_adapter(lora_id)
+
+ def list_loras(self) -> Set[int]:
+ if not self.lora_manager:
+ raise RuntimeError("LoRA is not enabled.")
+ return self.lora_manager.list_adapters()
+
+ def remove_all_prompt_adapters(self):
+ if not self.prompt_adapter_manager:
+ raise RuntimeError("PromptAdapter is not enabled.")
+ self.prompt_adapter_manager.remove_all_adapters()
+
+ def set_active_prompt_adapters(
+ self, prompt_adapter_requests: Set[PromptAdapterRequest],
+ prompt_adapter_mapping: PromptAdapterMapping) -> None:
+ if not self.prompt_adapter_manager:
+ raise RuntimeError("PromptAdapter is not enabled.")
+ self.prompt_adapter_manager.set_active_adapters(
+ prompt_adapter_requests, prompt_adapter_mapping)
+
+ def add_prompt_adapter(
+ self, prompt_adapter_request: PromptAdapterRequest) -> bool:
+ if not self.prompt_adapter_manager:
+ raise RuntimeError("PromptAdapter is not enabled.")
+ return self.prompt_adapter_manager.add_adapter(prompt_adapter_request)
+
+ def remove_prompt_adapter(self, prompt_adapter_id: int) -> bool:
+ if not self.prompt_adapter_manager:
+ raise RuntimeError("PromptAdapter is not enabled.")
+ return self.prompt_adapter_manager.remove_adapter(prompt_adapter_id)
+
+ def pin_prompt_adapter(self, prompt_adapter_id: int) -> bool:
+ if not self.prompt_adapter_manager:
+ raise RuntimeError("PromptAdapter is not enabled.")
+ return self.prompt_adapter_manager.pin_adapter(prompt_adapter_id)
+
+ def list_prompt_adapters(self) -> Set[int]:
+ if not self.prompt_adapter_manager:
+ raise RuntimeError("PromptAdapter is not enabled.")
+ return self.prompt_adapter_manager.list_adapters()
+
+ @property
+ def model_is_mrope(self) -> bool:
+ """Detect if the model has "mrope" rope_scaling type.
+ mrope requires keep "rope_deltas" between prompt and decoding phases."""
+ rope_scaling = getattr(self.model_config.hf_config, "rope_scaling", {})
+ if rope_scaling is None:
+ return False
+ return rope_scaling.get("type", None) == "mrope"
+
+ @torch.inference_mode()
+ def capture_model(self, kv_caches: List[List[torch.Tensor]]) -> None:
+ """Cuda graph capture a model.
+
+ Note that CUDA graph's performance gain is negligible if number
+ of batched tokens are larger than 200. And since CUDA graph
+ requires fixed sized tensors, supporting large/variable batch
+ size requires high GPU memory overhead. Thus, vLLM only captures
+ decoding requests. Mixed batch (chunked prefill + decoding) or
+ prefill requests are not captured.
+
+ Since it is used for decoding-only, it assumes there's only 1 token
+ per sequence in the batch.
+ """
+ assert not self.model_config.enforce_eager
+ logger.info("Capturing the model for CUDA graphs. This may lead to "
+ "unexpected consequences if the model is not static. To "
+ "run the model in eager mode, set 'enforce_eager=True' or "
+ "use '--enforce-eager' in the CLI.")
+ logger.info("CUDA graphs can take additional 1~3 GiB memory per GPU. "
+ "If you are running out of memory, consider decreasing "
+ "`gpu_memory_utilization` or enforcing eager mode. "
+ "You can also reduce the `max_num_seqs` as needed "
+ "to decrease memory usage.")
+ start_time = time.perf_counter()
+
+ # Prepare dummy inputs. These will be reused for all batch sizes.
+ max_batch_size = self.max_batchsize_to_capture
+ input_tokens = torch.zeros(max_batch_size, dtype=torch.long).cuda()
+ input_positions = torch.zeros(max_batch_size, dtype=torch.long).cuda()
+ if self.model_is_mrope:
+ input_positions = torch.tile(input_positions, (3, 1))
+ # Prepare dummy previous_hidden_states only if needed by the model.
+ # This is used by draft models such as EAGLE.
+ previous_hidden_states = None
+ if "previous_hidden_states" in inspect.signature(
+ self.model.forward).parameters:
+ previous_hidden_states = torch.empty(
+ [max_batch_size,
+ self.model_config.get_hidden_size()],
+ dtype=self.model_config.dtype,
+ device=self.device)
+
+ intermediate_inputs = None
+ if not get_pp_group().is_first_rank:
+ intermediate_inputs = self.model.make_empty_intermediate_tensors(
+ batch_size=max_batch_size,
+ dtype=self.model_config.dtype,
+ device=self.device)
+
+ # Prepare buffer for outputs. These will be reused for all batch sizes.
+ # It will be filled after the first graph capture.
+ hidden_or_intermediate_states: List[Optional[torch.Tensor]] = [
+ None
+ ] * self.parallel_config.pipeline_parallel_size
+
+ graph_batch_size = self.max_batchsize_to_capture
+ batch_size_capture_list = [
+ bs for bs in _BATCH_SIZES_TO_CAPTURE if bs <= graph_batch_size
+ ]
+
+ with self.attn_state.graph_capture(
+ max_batch_size), graph_capture() as graph_capture_context:
+ # NOTE: Capturing the largest batch size first may help reduce the
+ # memory usage of CUDA graph.
+ for virtual_engine in range(
+ self.parallel_config.pipeline_parallel_size):
+ for batch_size in reversed(batch_size_capture_list):
+ attn_metadata = (
+ self.attn_state.graph_capture_get_metadata_for_batch(
+ batch_size,
+ is_encoder_decoder_model=self.model_config.
+ is_encoder_decoder_model))
+
+ if self.lora_config:
+ lora_mapping = LoRAMapping(
+ **dict(index_mapping=[0] * batch_size,
+ prompt_mapping=[0] * batch_size,
+ is_prefill=False))
+ self.set_active_loras(set(), lora_mapping)
+
+ if self.prompt_adapter_config:
+ prompt_adapter_mapping = PromptAdapterMapping(
+ [-1] * batch_size,
+ [-1] * batch_size,
+ )
+ self.set_active_prompt_adapters(
+ set(), prompt_adapter_mapping)
+ graph_runner = CUDAGraphRunner(
+ self.model, self.attn_backend.get_name(),
+ self.attn_state.graph_clone(batch_size),
+ self.model_config.is_encoder_decoder_model)
+
+ capture_inputs = {
+ "input_ids":
+ input_tokens[:batch_size],
+ "positions":
+ input_positions[..., :batch_size],
+ "hidden_or_intermediate_states":
+ hidden_or_intermediate_states[
+ virtual_engine] # type: ignore
+ [:batch_size]
+ if hidden_or_intermediate_states[virtual_engine]
+ is not None else None,
+ "intermediate_inputs":
+ intermediate_inputs[:batch_size]
+ if intermediate_inputs is not None else None,
+ "kv_caches":
+ kv_caches[virtual_engine],
+ "attn_metadata":
+ attn_metadata,
+ "memory_pool":
+ self.graph_memory_pool,
+ "stream":
+ graph_capture_context.stream
+ }
+ if previous_hidden_states is not None:
+ capture_inputs[
+ "previous_hidden_states"] = previous_hidden_states[:
+ batch_size]
+
+ if self.has_inner_state:
+ # Only used by Mamba-based models CUDA graph atm (Jamba)
+ capture_inputs.update({
+ "seqlen_agnostic_capture_inputs":
+ self.model.get_seqlen_agnostic_capture_inputs(
+ batch_size)
+ })
+ if self.model_config.is_encoder_decoder_model:
+ # add the additional inputs to capture for
+ # encoder-decoder models.
+ self._update_inputs_to_capture_for_enc_dec_model(
+ capture_inputs)
+
+ with set_forward_context(attn_metadata):
+ graph_runner.capture(**capture_inputs)
+ self.graph_memory_pool = graph_runner.graph.pool()
+ self.graph_runners[virtual_engine][batch_size] = (
+ graph_runner)
+
+ end_time = time.perf_counter()
+ elapsed_time = end_time - start_time
+ # This usually takes < 10 seconds.
+ logger.info("Graph capturing finished in %.0f secs.", elapsed_time)
+
+ def _update_inputs_to_capture_for_enc_dec_model(self,
+ capture_inputs: Dict[str,
+ Any]):
+ """
+ Updates the set of input tensors needed for CUDA graph capture in an
+ encoder-decoder model.
+
+ This method modifies the provided `capture_inputs` dictionary by
+ adding tensors specific to encoder-decoder specific models that
+ need to be captured for CUDA Graph replay.
+ """
+ # During the decode phase encoder_input_ids and encoder_positions are
+ # unset. Do the same thing for graph capture.
+ capture_inputs["encoder_input_ids"] = torch.tensor(
+ [], dtype=torch.long).cuda()
+ capture_inputs["encoder_positions"] = torch.tensor(
+ [], dtype=torch.long).cuda()
+
+ @property
+ def vocab_size(self) -> int:
+ return self.model_config.get_vocab_size()
+
+
+class ModelRunner(GPUModelRunnerBase[ModelInputForGPUWithSamplingMetadata]):
+ """
+ GPU model runner with sampling step.
+ """
+ _model_input_cls: Type[ModelInputForGPUWithSamplingMetadata] = (
+ ModelInputForGPUWithSamplingMetadata)
+ _builder_cls: Type[ModelInputForGPUBuilder] = ModelInputForGPUBuilder
+
+ def make_model_input_from_broadcasted_tensor_dict(
+ self,
+ tensor_dict: Dict[str, Any],
+ ) -> ModelInputForGPUWithSamplingMetadata:
+ model_input = \
+ ModelInputForGPUWithSamplingMetadata.from_broadcasted_tensor_dict(
+ tensor_dict,
+ attn_backend=self.attn_backend,
+ )
+ return model_input
+
+ def prepare_model_input(
+ self,
+ seq_group_metadata_list: List[SequenceGroupMetadata],
+ virtual_engine: int = 0,
+ finished_requests_ids: Optional[List[str]] = None,
+ ) -> ModelInputForGPUWithSamplingMetadata:
+ """Prepare the model input based on a given sequence group, including
+ metadata for the sampling step.
+
+ The API assumes seq_group_metadata_list is sorted by prefill -> decode.
+
+ The result tensors and data structure also batches input in prefill
+ -> decode order. For example,
+
+ - input_tokens[:num_prefill_tokens] contains prefill tokens.
+ - input_tokens[num_prefill_tokens:] contains decode tokens.
+
+ If cuda graph is required, this API automatically pads inputs.
+ """
+ model_input = self._prepare_model_input_tensors(
+ seq_group_metadata_list, finished_requests_ids)
+ if get_pp_group().is_last_rank:
+ # Sampling metadata is only required for the final pp group
+ generators = self.get_generators(finished_requests_ids)
+ sampling_metadata = SamplingMetadata.prepare(
+ seq_group_metadata_list, model_input.seq_lens,
+ model_input.query_lens, self.device, self.pin_memory,
+ generators, self.sampling_metadata_cache)
+ else:
+ sampling_metadata = None
+ is_prompt = (seq_group_metadata_list[0].is_prompt
+ if seq_group_metadata_list else None)
+ return dataclasses.replace(model_input,
+ sampling_metadata=sampling_metadata,
+ is_prompt=is_prompt,
+ virtual_engine=virtual_engine)
+
+ @torch.inference_mode()
+ # @dump_input_when_exception(exclude_args=[0], exclude_kwargs=["self"])
+ def execute_model(
+ self,
+ model_input: ModelInputForGPUWithSamplingMetadata,
+ kv_caches: List[torch.Tensor],
+ intermediate_tensors: Optional[IntermediateTensors] = None,
+ num_steps: int = 1,
+ ) -> Optional[Union[List[SamplerOutput], IntermediateTensors]]:
+ if num_steps > 1:
+ raise ValueError("num_steps > 1 is not supported in ModelRunner")
+
+ if self.lora_config:
+ assert model_input.lora_requests is not None
+ assert model_input.lora_mapping is not None
+ self.set_active_loras(model_input.lora_requests,
+ model_input.lora_mapping)
+
+ if self.prompt_adapter_config:
+ assert model_input.prompt_adapter_requests is not None
+ assert model_input.prompt_adapter_mapping is not None
+ self.set_active_prompt_adapters(
+ model_input.prompt_adapter_requests,
+ model_input.prompt_adapter_mapping)
+
+ self.attn_state.begin_forward(model_input)
+
+ # Currently cuda graph is only supported by the decode phase.
+ assert model_input.attn_metadata is not None
+ prefill_meta = model_input.attn_metadata.prefill_metadata
+ decode_meta = model_input.attn_metadata.decode_metadata
+ # TODO(andoorve): We can remove this once all
+ # virtual engines share the same kv cache.
+ virtual_engine = model_input.virtual_engine
+ if prefill_meta is None and decode_meta.use_cuda_graph:
+ assert model_input.input_tokens is not None
+ graph_batch_size = model_input.input_tokens.shape[0]
+ model_executable = self.graph_runners[virtual_engine][
+ graph_batch_size]
+ else:
+ model_executable = self.model
+
+ multi_modal_kwargs = model_input.multi_modal_kwargs or {}
+ seqlen_agnostic_kwargs = {
+ "finished_requests_ids": model_input.finished_requests_ids,
+ "request_ids_to_seq_ids": model_input.request_ids_to_seq_ids,
+ } if self.has_inner_state else {}
+ if (self.observability_config is not None
+ and self.observability_config.collect_model_forward_time):
+ model_forward_start = torch.cuda.Event(enable_timing=True)
+ model_forward_end = torch.cuda.Event(enable_timing=True)
+ model_forward_start.record()
+
+ with set_forward_context(model_input.attn_metadata):
+ hidden_or_intermediate_states = model_executable(
+ input_ids=model_input.input_tokens,
+ positions=model_input.input_positions,
+ kv_caches=kv_caches,
+ attn_metadata=model_input.attn_metadata,
+ intermediate_tensors=intermediate_tensors,
+ **MultiModalInputs.as_kwargs(multi_modal_kwargs,
+ device=self.device),
+ **seqlen_agnostic_kwargs)
+
+ if (self.observability_config is not None
+ and self.observability_config.collect_model_forward_time):
+ model_forward_end.record()
+
+ # Compute the logits in the last pipeline stage.
+ if not get_pp_group().is_last_rank:
+ if (self.is_driver_worker
+ and hidden_or_intermediate_states is not None
+ and isinstance(hidden_or_intermediate_states,
+ IntermediateTensors)
+ and self.observability_config is not None
+ and self.observability_config.collect_model_forward_time):
+ model_forward_end.synchronize()
+ model_forward_time = model_forward_start.elapsed_time(
+ model_forward_end)
+ orig_model_forward_time = 0.0
+ if intermediate_tensors is not None:
+ orig_model_forward_time = intermediate_tensors.tensors.get(
+ "model_forward_time", torch.tensor(0.0)).item()
+ hidden_or_intermediate_states.tensors["model_forward_time"] = (
+ torch.tensor(model_forward_time + orig_model_forward_time))
+ return hidden_or_intermediate_states
+
+ logits = self.model.compute_logits(hidden_or_intermediate_states,
+ model_input.sampling_metadata)
+
+ if not self.is_driver_worker:
+ return []
+
+ if model_input.async_callback is not None:
+ model_input.async_callback()
+
+ # Sample the next token.
+ output: SamplerOutput = self.model.sample(
+ logits=logits,
+ sampling_metadata=model_input.sampling_metadata,
+ )
+ if (self.observability_config is not None
+ and self.observability_config.collect_model_forward_time
+ and output is not None):
+ model_forward_end.synchronize()
+ model_forward_time = model_forward_start.elapsed_time(
+ model_forward_end)
+ orig_model_forward_time = 0.0
+ if intermediate_tensors is not None:
+ orig_model_forward_time = intermediate_tensors.tensors.get(
+ "model_forward_time", torch.tensor(0.0)).item()
+ # If there are multiple workers, we are still tracking the latency
+ # from the start time of the driver worker to the end time of the
+ # driver worker. The model forward time will then end up covering
+ # the communication time as well.
+ output.model_forward_time = (orig_model_forward_time +
+ model_forward_time)
+
+ if self.return_hidden_states:
+ # we only need to pass hidden states of most recent token
+ assert model_input.sampling_metadata is not None
+ indices = model_input.sampling_metadata.selected_token_indices
+ if model_input.is_prompt:
+ hidden_states = hidden_or_intermediate_states.index_select(
+ 0, indices)
+ output.prefill_hidden_states = hidden_or_intermediate_states
+ elif decode_meta.use_cuda_graph:
+ hidden_states = hidden_or_intermediate_states[:len(indices)]
+ else:
+ hidden_states = hidden_or_intermediate_states
+
+ output.hidden_states = hidden_states
+
+ return [output]
+
+
+class CUDAGraphRunner:
+
+ def __init__(self, model: nn.Module, backend_name: str,
+ attn_state: AttentionState, is_encoder_decoder_model: bool):
+ self.model = model
+ self.backend_name = backend_name
+ self.attn_state = attn_state
+
+ self.input_buffers: Dict[str, torch.Tensor] = {}
+ self.output_buffers: Dict[str, torch.Tensor] = {}
+
+ self._graph: Optional[torch.cuda.CUDAGraph] = None
+ self._is_encoder_decoder_model = is_encoder_decoder_model
+
+ @property
+ def graph(self):
+ assert self._graph is not None
+ return self._graph
+
+ def capture(
+ self,
+ input_ids: torch.Tensor,
+ positions: torch.Tensor,
+ hidden_or_intermediate_states: Optional[Union[IntermediateTensors,
+ torch.Tensor]],
+ intermediate_inputs: Optional[IntermediateTensors],
+ kv_caches: List[torch.Tensor],
+ attn_metadata: AttentionMetadata,
+ memory_pool: Optional[Tuple[int, int]],
+ stream: torch.cuda.Stream,
+ **kwargs,
+ ) -> Union[torch.Tensor, IntermediateTensors]:
+ assert self._graph is None
+ # Run the model a few times without capturing the graph.
+ # This is to make sure that the captured graph does not include the
+ # kernel launches for initial benchmarking (e.g., Triton autotune).
+ # Note one iteration is not enough for torch.jit.script
+ for _ in range(_NUM_WARMUP_ITERS):
+ self.model(
+ input_ids=input_ids,
+ positions=positions,
+ kv_caches=kv_caches,
+ attn_metadata=attn_metadata,
+ intermediate_tensors=intermediate_inputs,
+ **kwargs,
+ )
+ # Wait for the warm up operations to finish before proceeding with
+ # Graph Capture.
+ torch.cuda.synchronize()
+ # Capture the graph.
+ self._graph = torch.cuda.CUDAGraph()
+ with torch.cuda.graph(self._graph, pool=memory_pool, stream=stream):
+ output_hidden_or_intermediate_states = self.model(
+ input_ids=input_ids,
+ positions=positions,
+ kv_caches=kv_caches,
+ attn_metadata=attn_metadata,
+ intermediate_tensors=intermediate_inputs,
+ **kwargs,
+ )
+ if hidden_or_intermediate_states is not None:
+ if get_pp_group().is_last_rank:
+ hidden_or_intermediate_states.copy_(
+ output_hidden_or_intermediate_states)
+ else:
+ for key in hidden_or_intermediate_states.tensors:
+ hidden_or_intermediate_states[key].copy_(
+ output_hidden_or_intermediate_states[key])
+ else:
+ hidden_or_intermediate_states = (
+ output_hidden_or_intermediate_states)
+
+ del output_hidden_or_intermediate_states
+ # make sure `output_hidden_states` is deleted
+ # in the graph's memory pool
+ gc.collect()
+ torch.cuda.synchronize()
+
+ # Save the input and output buffers.
+ self.input_buffers = {
+ "input_ids":
+ input_ids,
+ "positions":
+ positions,
+ "kv_caches":
+ kv_caches,
+ **self.attn_state.get_graph_input_buffers(
+ attn_metadata, self._is_encoder_decoder_model),
+ **kwargs,
+ }
+ if intermediate_inputs is not None:
+ self.input_buffers.update(intermediate_inputs.tensors)
+ if get_pp_group().is_last_rank:
+ self.output_buffers = {
+ "hidden_states": hidden_or_intermediate_states
+ }
+ else:
+ self.output_buffers = hidden_or_intermediate_states
+ return hidden_or_intermediate_states
+
+ def forward(
+ self,
+ input_ids: torch.Tensor,
+ positions: torch.Tensor,
+ kv_caches: List[torch.Tensor],
+ attn_metadata: AttentionMetadata,
+ intermediate_tensors: Optional[IntermediateTensors],
+ **kwargs,
+ ) -> torch.Tensor:
+ # KV caches are fixed tensors, so we don't need to copy them.
+ del kv_caches
+
+ # Copy the input tensors to the input buffers.
+ self.input_buffers["input_ids"].copy_(input_ids, non_blocking=True)
+ self.input_buffers["positions"].copy_(positions, non_blocking=True)
+
+ if self.backend_name != "placeholder-attn":
+ self.input_buffers["slot_mapping"].copy_(
+ attn_metadata.slot_mapping, non_blocking=True)
+
+ self.attn_state.prepare_graph_input_buffers(
+ self.input_buffers, attn_metadata, self._is_encoder_decoder_model)
+
+ if "seqlen_agnostic_capture_inputs" in self.input_buffers:
+ self.model.copy_inputs_before_cuda_graphs(self.input_buffers,
+ **kwargs)
+
+ if "previous_hidden_states" in self.input_buffers:
+ self.input_buffers["previous_hidden_states"].copy_(
+ kwargs["previous_hidden_states"], non_blocking=True)
+
+ if intermediate_tensors is not None:
+ for key in intermediate_tensors.tensors:
+ if key != "model_execute_time" and key != "model_forward_time":
+ self.input_buffers[key].copy_(intermediate_tensors[key],
+ non_blocking=True)
+ if self._is_encoder_decoder_model:
+ self.input_buffers["encoder_input_ids"].copy_(
+ kwargs['encoder_input_ids'], non_blocking=True)
+ self.input_buffers["encoder_positions"].copy_(
+ kwargs['encoder_positions'], non_blocking=True)
+
+ # Run the graph.
+ self.graph.replay()
+ # Return the output tensor.
+ if get_pp_group().is_last_rank:
+ return self.output_buffers["hidden_states"]
+
+ return self.output_buffers
+
+ def __call__(self, *args, **kwargs):
+ return self.forward(*args, **kwargs)
+
+
+def _get_graph_batch_size(batch_size: int) -> int:
+ """Returns the padded batch size given actual batch size.
+
+ Batch sizes are 1, 2, 4, _BATCH_SIZE_ALIGNMENT,
+ 2*_BATCH_SIZE_ALIGNMENT, 3*_BATCH_SIZE_ALIGNMENT...
+ """
+ if batch_size <= 2:
+ return batch_size
+ elif batch_size <= 4:
+ return 4
+ else:
+ return ((batch_size + _BATCH_SIZE_ALIGNMENT - 1) //
+ _BATCH_SIZE_ALIGNMENT * _BATCH_SIZE_ALIGNMENT)
+
+
+def _get_max_graph_batch_size(max_num_seqs: int) -> int:
+ """
+ max_num_seqs: Maximum number of sequences in a batch.
+ _BATCH_SIZES_TO_CAPTURE: all the sizes that we want to capture.
+
+ pad the max_num_seqs if necessary by calling _get_graph_batch_size,
+ which will deal with some edge cases like 1, 2, 4.
+
+ if the padded size is in _BATCH_SIZES_TO_CAPTURE, return the padded size.
+ if not, it means the padded size is larger than the largest size in
+ _BATCH_SIZES_TO_CAPTURE, return the largest size in _BATCH_SIZES_TO_CAPTURE.
+ """
+ padded_size = _get_graph_batch_size(max_num_seqs)
+ if padded_size in _BATCH_SIZES_TO_CAPTURE:
+ return padded_size
+ assert padded_size > _BATCH_SIZES_TO_CAPTURE[-1]
+ return _BATCH_SIZES_TO_CAPTURE[-1]
diff --git a/qwen3_6_scripts/patch_ops.sh b/qwen3_6_scripts/patch_ops.sh
index 1035b2fe..cc0d0b3e 100755
--- a/qwen3_6_scripts/patch_ops.sh
+++ b/qwen3_6_scripts/patch_ops.sh
@@ -1,94 +1,96 @@
-# BI-V100 patch script for Qwen3.6-27B (Qwen3_5 architecture)
+# BI-V100 engine patches for Qwen3.6-35B-A3B (Qwen3_5 architecture)
#
-# Triton situation on BI-V100:
-# - Standard Triton 2.3.1 is already present in the image.
-# - HAS_TRITON = False (hardcoded in vendor vllm), but Triton is still used
-# for TP-mode cache management (custom_cache_manager / libentry).
-# - The vendor's triton_utils/__init__.py, custom_cache_manager.py, libentry.py
-# are already correct for standard Triton 2.3.1 — do NOT overwrite them.
-# - DO NOT install BI-V150 corex Triton 2.1.0 (pkgs/triton): that causes
-# GPU hang on BI-V100 because the Triton CUDA PTX kernels are incompatible.
-
-# Recommended server start command for TP=4 support 100K, need chunked prefill
-# CUDA_VISIBLE_DEVICES="4,5,6,7" VLLM_ENGINE_ITERATION_TIMEOUT_S=3600 python3 -m vllm.entrypoints.openai.api_server \
-# --model /workspace/models/Qwen3.6-27B --port 1111 --served-model-name llm \
-# --max-model-len 100000 --enforce-eager --trust-remote-code -tp 4 --gpu-memory-utilization 0.95 \
-# --max-num-seqs 1 --disable-log-requests --disable-frontend-multiprocessing \
-# --max-num-batched-tokens 4096 --enable-chunked-prefill
+# All modifications are FULL FILE REPLACEMENTS — no AST patch scripts.
+# Each file was read in full from the base image vllm source, modified
+# with the necessary fixes, and placed here as a complete copy.
#
-# With prefix caching (GDN align-mode, requires chunked prefill):
-# CUDA_VISIBLE_DEVICES="4,5,6,7" VLLM_ENGINE_ITERATION_TIMEOUT_S=3600 python3 -m vllm.entrypoints.openai.api_server \
-# --model /workspace/models/Qwen3.6-35B-A3B --port 1111 --served-model-name llm \
-# --max-model-len 150000 --trust-remote-code -tp 4 --gpu-memory-utilization 0.90 \
-# --max-num-seqs 1 --disable-log-requests --disable-frontend-multiprocessing \
-# --max-num-batched-tokens 8192 --enable-chunked-prefill --enable-prefix-caching \
-# --max-seq-len-to-capture 32768
+# Base image: git.modelhub.org.cn:9443/enginex-iluvatar/bi100-3.2.3-x86-ubuntu20.04-py3.10-poc-llm-infer:v1.2.3
+# vllm install path: /usr/local/corex/lib/python3/dist-packages/vllm/
-# --- paged_attn.py: replace forward_prefix with pure-PyTorch fallback -------
-# The Triton context_attention_fwd kernel hangs BI-V100 GPUs permanently
-# (standard Triton 2.3.1 PTX is not supported by the corex runtime either).
-# Our paged_attn.py bypasses it entirely via _forward_prefix_pytorch, which
-# utilizes K-tiling techniques, and also have _forward_decode_pytorch to bypass kernel
-# when context length is high
-cp ./paged_attn.py /usr/local/corex/lib/python3/dist-packages/vllm/attention/ops/paged_attn.py
+VLLM=/usr/local/corex/lib/python3/dist-packages/vllm
+VLLM64=/usr/local/corex/lib64/python3/dist-packages/vllm
-# --- model_runner.py: fix prefix_cache_hit stays True in chunked-prefill chunk 2+ ---
-# Bug: _compute_for_prefix_cache_hit Case 1 (prefix_cache_len <= context_len)
-# leaves prefix_cache_hit=True. Then _add_seq_group uses block_table=computed_block_nums
-# (only the original prefix blocks), ignoring chunk-1 KV cache blocks.
-# _forward_prefix_pytorch then gets an undersized block_tables and crashes with
-# "amax(): Expected reduction dim -1 to have non-zero size" on the 2nd tile.
-# Fix: set prefix_cache_hit=False for Case 1 so the full block_tables is used.
-python3 ./patch_model_runner.py
+# Detect which lib path exists
+if [ -d "$VLLM" ]; then
+ V=$VLLM
+elif [ -d "$VLLM64" ]; then
+ V=$VLLM64
+else
+ echo "[patch_ops] ERROR: vllm not found at lib or lib64 path"
+ exit 1
+fi
+
+echo "[patch_ops] vllm path: $V"
+
+# --- paged_attn.py: pure-PyTorch attention fallback --------------------------
+# Bypasses Triton context_attention_fwd (hangs BI-V100 permanently).
+# Uses K-tiling Flash Attention online softmax for prefix attention.
+# Uses pure-PyTorch decode for seq_len > 32K.
+# CCCL-ported: adaptive tile sizing from dispatch_reduce.cuh GridEvenShare.
+cp ./paged_attn.py $V/attention/ops/paged_attn.py
+echo "[patch_ops] paged_attn.py → attention/ops/"
+
+# --- model_runner.py: prefix_cache_hit fix -----------------------------------
+# Bug: Case 1 (prefix_cache_len <= context_len) leaves prefix_cache_hit=True,
+# causing undersized block_tables in chunked prefill chunk 2+.
+# Fix: set prefix_cache_hit=False for Case 1.
+# FULL FILE REPLACEMENT — no patch_model_runner.py script.
+cp ./model_runner.py $V/worker/model_runner.py
+echo "[patch_ops] model_runner.py → worker/"
+
+# --- xformers.py: head_dim>128 fallback + Q-tiling --------------------------
+# Injects _run_sdpa_fallback (pure matmul+softmax) for head_dim=256.
+# ixformer flash attention crashes (is_causal=True) or gives wrong output.
+# Also disables auto chunked-prefill (Q-tiling handles long context).
+# FULL FILE REPLACEMENT — no patch_xformers_sdpa_seq.py script.
+cp ./xformers.py $V/attention/backends/xformers.py
+echo "[patch_ops] xformers.py → attention/backends/"
+
+# --- arg_utils.py: disable auto chunked-prefill for 32K+ --------------------
+# Q-tiling in _run_sdpa_fallback handles long-context memory.
+# FULL FILE REPLACEMENT.
+cp ./arg_utils.py $V/engine/arg_utils.py
+echo "[patch_ops] arg_utils.py → engine/"
+
+# --- logits_processor.py: seq_groups=None guard ------------------------------
+# Prevents crash when seq_groups is None during intermediate chunked-prefill.
+# FULL FILE REPLACEMENT.
+cp ./logits_processor.py $V/model_executor/layers/logits_processor.py
+echo "[patch_ops] logits_processor.py → model_executor/layers/"
# --- transformers: Qwen3_5 tokenizer / model files --------------------------
pip install transformers==4.55.3 -i https://pypi.tuna.tsinghua.edu.cn/simple
cp -r ./qwen3_5 /usr/local/lib/python3.10/site-packages/transformers/models/
cp -r ./qwen3_5_moe /usr/local/lib/python3.10/site-packages/transformers/models/
python3 ./patch_transformers_qwen3_5.py
+echo "[patch_ops] transformers Qwen3_5 models installed"
-# --- vllm model: Qwen3.6-27B (Qwen3_5 arch) --------------------------------
-cp ./mamba_cache.py /usr/local/corex/lib/python3/dist-packages/vllm/model_executor/models/
-cp ./qwen3_5.py /usr/local/corex/lib/python3/dist-packages/vllm/model_executor/models/qwen3_5.py
+# --- vllm model: Qwen3.6 (Qwen3_5 arch) ------------------------------------
+cp ./mamba_cache.py $V/model_executor/models/
+cp ./qwen3_5.py $V/model_executor/models/qwen3_5.py
python3 ./patch_vllm_qwen3_5.py
+echo "[patch_ops] qwen3_5.py model registered"
-# --- sequence.py: fix completion_tokens inflation under chunked prefill ------
-# Bug: get_output_token_ids_to_return(delta=True) with num_new_tokens=0
-# returns _cached_all_token_ids[-0:] == [0:] (the ENTIRE prompt+output list).
-# Each prefill chunk step adds prompt_len to previous_num_tokens, so a 10K
-# prompt processed in 3 chunks inflates completion_tokens by ~30K.
-# Also adds num_cached_tokens field to RequestMetrics for prefix-cache stats.
-cp ./sequence.py /usr/local/corex/lib/python3/dist-packages/vllm/sequence.py
+# --- sequence.py: fix completion_tokens inflation ----------------------------
+cp ./sequence.py $V/sequence.py
+echo "[patch_ops] sequence.py → /"
-# --- scheduler.py: record num_cached_tokens in RequestMetrics ----------------
-# Sets seq_group.metrics.num_cached_tokens = prefix_cache_len on first prefill
-# when --enable-prefix-caching is active, so serving_chat.py can report it in
-# usage.prompt_tokens_details.cached_tokens (OpenAI-compatible API response).
-cp ./scheduler.py /usr/local/corex/lib/python3/dist-packages/vllm/core/scheduler.py
+# --- scheduler.py: record num_cached_tokens ---------------------------------
+cp ./scheduler.py $V/core/scheduler.py
+echo "[patch_ops] scheduler.py → core/"
-# --- xformers: bypass cudnnFlashAttnForward (head_dim=256 > 128 limit) ------
-# Injects _run_sdpa_fallback (pure matmul+softmax) into xformers.py.
-# Required because head_dim=256 > 128 and ixformer flash attention either
-# crashes (is_causal=True) or produces wrong output (attn_mask path).
-# The fallback uses query_start_loc to derive actual query lengths, so it
-# works correctly during profiling runs with chunked-prefill-style batches.
-# also bypasses auto chunked prefill on
-python3 ./patch_xformers_sdpa_seq.py
-
-# --- tool parser: Qwen3 XML tool call format ---------------------------------
-# Registers "qwen3_coder" parser for Qwen3.6 XML-style tool calls:
-# \nvalue\n
-# Use at server start: --tool-call-parser qwen3_coder --enable-auto-tool-choice
-cp ./qwen3coder_tool_parser.py /usr/local/corex/lib/python3/dist-packages/vllm/entrypoints/openai/tool_parsers/
+# --- tool parser: Qwen3 XML tool call format --------------------------------
+cp ./qwen3coder_tool_parser.py $V/entrypoints/openai/tool_parsers/
python3 ./patch_vllm_tool_parser.py
+echo "[patch_ops] qwen3_coder tool parser registered"
-# --- reasoning parser: Qwen3 ... split ------------------------
-# Adds --reasoning-parser qwen3 support.
-# Routes thinking tokens to reasoning_content, rest to content in the delta.
-# Works together with --tool-call-parser qwen3_coder (think → tool call flow).
-cp -r ./reasoning /usr/local/corex/lib/python3/dist-packages/vllm/
-cp ./protocol.py /usr/local/corex/lib/python3/dist-packages/vllm/entrypoints/openai/protocol.py
-cp ./cli_args.py /usr/local/corex/lib/python3/dist-packages/vllm/entrypoints/openai/cli_args.py
-cp ./serving_chat.py /usr/local/corex/lib/python3/dist-packages/vllm/entrypoints/openai/serving_chat.py
-cp ./api_server.py /usr/local/corex/lib/python3/dist-packages/vllm/entrypoints/openai/api_server.py
-cp ./chat_utils.py /usr/local/corex/lib/python3/dist-packages/vllm/entrypoints/chat_utils.py
+# --- reasoning parser: Qwen3 ... split -----------------------
+cp -r ./reasoning $V/
+cp ./protocol.py $V/entrypoints/openai/protocol.py
+cp ./cli_args.py $V/entrypoints/openai/cli_args.py
+cp ./serving_chat.py $V/entrypoints/openai/serving_chat.py
+cp ./api_server.py $V/entrypoints/openai/api_server.py
+cp ./chat_utils.py $V/entrypoints/chat_utils.py
+echo "[patch_ops] reasoning parser + serving files installed"
+
+echo "[patch_ops] DONE — all patches applied via full file replacement"
diff --git a/qwen3_6_scripts/xformers.py b/qwen3_6_scripts/xformers.py
new file mode 100644
index 00000000..c0bd39af
--- /dev/null
+++ b/qwen3_6_scripts/xformers.py
@@ -0,0 +1,907 @@
+"""Attention layer with xFormers and PagedAttention."""
+from dataclasses import dataclass
+from typing import Any, Dict, List, Optional, Tuple, Type
+
+import torch
+# from xformers import ops as xops
+from ixformer.contrib.xformers import ops as xops
+from xformers.ops.fmha.attn_bias import (AttentionBias,
+ BlockDiagonalMask,)
+from ixformer.contrib.xformers.ops.fmha.attn_bias import (BlockDiagonalCausalMask,
+ LowerTriangularMaskWithTensorBias)
+
+from vllm.attention.backends.abstract import (AttentionBackend, AttentionImpl,
+ AttentionMetadata, AttentionType)
+from vllm.attention.backends.utils import (CommonAttentionState,
+ CommonMetadataBuilder)
+from vllm.attention.ops.paged_attn import (PagedAttention,
+ PagedAttentionMetadata)
+from vllm.logger import init_logger
+
+logger = init_logger(__name__)
+
+
+class XFormersBackend(AttentionBackend):
+
+ @staticmethod
+ def get_name() -> str:
+ return "xformers"
+
+ @staticmethod
+ def get_impl_cls() -> Type["XFormersImpl"]:
+ return XFormersImpl
+
+ @staticmethod
+ def get_metadata_cls() -> Type["AttentionMetadata"]:
+ return XFormersMetadata
+
+ @staticmethod
+ def get_builder_cls() -> Type["XFormersMetadataBuilder"]:
+ return XFormersMetadataBuilder
+
+ @staticmethod
+ def get_state_cls() -> Type["CommonAttentionState"]:
+ return CommonAttentionState
+
+ @staticmethod
+ def get_kv_cache_shape(
+ num_blocks: int,
+ block_size: int,
+ num_kv_heads: int,
+ head_size: int,
+ ) -> Tuple[int, ...]:
+ return PagedAttention.get_kv_cache_shape(num_blocks, block_size,
+ num_kv_heads, head_size)
+
+ @staticmethod
+ def swap_blocks(
+ src_kv_cache: torch.Tensor,
+ dst_kv_cache: torch.Tensor,
+ src_to_dst: Dict[int, int],
+ ) -> None:
+ PagedAttention.swap_blocks(src_kv_cache, dst_kv_cache, src_to_dst)
+
+ @staticmethod
+ def copy_blocks(
+ kv_caches: List[torch.Tensor],
+ src_to_dists: torch.Tensor,
+ ) -> None:
+ PagedAttention.copy_blocks(kv_caches, src_to_dists)
+
+
+@dataclass
+class XFormersMetadata(AttentionMetadata, PagedAttentionMetadata):
+ """Metadata for XFormersbackend.
+
+ NOTE: Any python object stored here is not updated when it is
+ cuda-graph replayed. If you have values that need to be changed
+ dynamically, it should be stored in tensor. The tensor has to be
+ updated from `CUDAGraphRunner.forward` API.
+ """
+
+ # |---------- N-1 iteration --------|
+ # |---------------- N iteration ---------------------|
+ # |- tokenA -|......................|-- newTokens ---|
+ # |---------- context_len ----------|
+ # |-------------------- seq_len ----------------------|
+ # |-- query_len ---|
+
+ # seq_lens stored as a tensor.
+ seq_lens_tensor: Optional[torch.Tensor]
+
+ # FIXME: It is for flash attn.
+ # Maximum sequence length among prefill batch. 0 if there are decoding
+ # requests only.
+ max_prefill_seq_len: int
+ # Maximum sequence length among decode batch. 0 if there are prefill
+ # requests only.
+ max_decode_seq_len: int
+
+ # Whether or not if cuda graph is enabled.
+ # Cuda-graph is currently enabled for decoding only.
+ # TODO(woosuk): Move `use_cuda_graph` out since it's unrelated to attention.
+ use_cuda_graph: bool
+
+ # (batch_size,). The sequence length per sequence. Sequence length means
+ # the computed tokens + new tokens None if it is a decoding.
+ seq_lens: Optional[List[int]] = None
+
+ # FIXME: It is for flash attn.
+ # (batch_size + 1,). The cumulative sequence lengths of the sequences in
+ # the batch, used to index into sequence. E.g., if the sequence length is
+ # [4, 6], it is [0, 4, 10].
+ seq_start_loc: Optional[torch.Tensor] = None
+
+ # (batch_size,) A tensor of context lengths (tokens that are computed
+ # so far).
+ context_lens_tensor: Optional[torch.Tensor] = None
+
+ # Maximum query length in the batch. None for decoding.
+ max_query_len: Optional[int] = None
+
+ # Max number of query tokens among request in the batch.
+ max_decode_query_len: Optional[int] = None
+
+ # (batch_size + 1,). The cumulative subquery lengths of the sequences in
+ # the batch, used to index into subquery. E.g., if the subquery length
+ # is [4, 6], it is [0, 4, 10].
+ query_start_loc: Optional[torch.Tensor] = None
+
+ # Self-attention prefill/decode metadata cache
+ _cached_prefill_metadata: Optional["XFormersMetadata"] = None
+ _cached_decode_metadata: Optional["XFormersMetadata"] = None
+
+ # Begin encoder attn & enc/dec cross-attn fields...
+
+ # Encoder sequence lengths representation
+ encoder_seq_lens: Optional[List[int]] = None
+ encoder_seq_lens_tensor: Optional[torch.Tensor] = None
+
+ # Maximum sequence length among encoder sequences
+ max_encoder_seq_len: Optional[int] = None
+
+ # Number of tokens input to encoder
+ num_encoder_tokens: Optional[int] = None
+
+ # Cross-attention memory-mapping data structures: slot mapping
+ # and block tables
+ cross_slot_mapping: Optional[torch.Tensor] = None
+ cross_block_tables: Optional[torch.Tensor] = None
+
+ def __post_init__(self):
+ # Set during the execution of the first attention op.
+ # It is a list because it is needed to set per prompt
+ # when alibi slopes is used. It is because of the limitation
+ # from xformer API.
+ # will not appear in the __repr__ and __init__
+ self.attn_bias: Optional[List[AttentionBias]] = None
+ self.encoder_attn_bias: Optional[List[AttentionBias]] = None
+ self.cross_attn_bias: Optional[List[AttentionBias]] = None
+
+ @property
+ def is_all_encoder_attn_metadata_set(self):
+ '''
+ All attention metadata required for encoder attention is set.
+ '''
+ return ((self.encoder_seq_lens is not None)
+ and (self.encoder_seq_lens_tensor is not None)
+ and (self.max_encoder_seq_len is not None))
+
+ @property
+ def is_all_cross_attn_metadata_set(self):
+ '''
+ All attention metadata required for enc/dec cross-attention is set.
+
+ Superset of encoder attention required metadata.
+ '''
+ return (self.is_all_encoder_attn_metadata_set
+ and (self.cross_slot_mapping is not None)
+ and (self.cross_block_tables is not None))
+
+ @property
+ def prefill_metadata(self) -> Optional["XFormersMetadata"]:
+ if self.num_prefills == 0:
+ return None
+
+ if self._cached_prefill_metadata is not None:
+ # Recover cached prefill-phase attention
+ # metadata structure
+ return self._cached_prefill_metadata
+
+ assert ((self.seq_lens is not None)
+ or (self.encoder_seq_lens is not None))
+ assert ((self.seq_lens_tensor is not None)
+ or (self.encoder_seq_lens_tensor is not None))
+
+ # Compute some attn_metadata fields which default to None
+ query_start_loc = (None if self.query_start_loc is None else
+ self.query_start_loc[:self.num_prefills + 1])
+ slot_mapping = (None if self.slot_mapping is None else
+ self.slot_mapping[:self.num_prefill_tokens])
+ seq_lens = (None if self.seq_lens is None else
+ self.seq_lens[:self.num_prefills])
+ seq_lens_tensor = (None if self.seq_lens_tensor is None else
+ self.seq_lens_tensor[:self.num_prefills])
+ context_lens_tensor = (None if self.context_lens_tensor is None else
+ self.context_lens_tensor[:self.num_prefills])
+ block_tables = (None if self.block_tables is None else
+ self.block_tables[:self.num_prefills])
+
+ # Construct & cache prefill-phase attention metadata structure
+ self._cached_prefill_metadata = XFormersMetadata(
+ num_prefills=self.num_prefills,
+ num_prefill_tokens=self.num_prefill_tokens,
+ num_decode_tokens=0,
+ slot_mapping=slot_mapping,
+ seq_lens=seq_lens,
+ seq_lens_tensor=seq_lens_tensor,
+ max_query_len=self.max_query_len,
+ max_prefill_seq_len=self.max_prefill_seq_len,
+ max_decode_seq_len=0,
+ query_start_loc=query_start_loc,
+ context_lens_tensor=context_lens_tensor,
+ block_tables=block_tables,
+ use_cuda_graph=False,
+ # Begin encoder & cross attn fields below...
+ encoder_seq_lens=self.encoder_seq_lens,
+ encoder_seq_lens_tensor=self.encoder_seq_lens_tensor,
+ max_encoder_seq_len=self.max_encoder_seq_len,
+ cross_slot_mapping=self.cross_slot_mapping,
+ cross_block_tables=self.cross_block_tables)
+ return self._cached_prefill_metadata
+
+ @property
+ def decode_metadata(self) -> Optional["XFormersMetadata"]:
+ if self.num_decode_tokens == 0:
+ return None
+
+ if self._cached_decode_metadata is not None:
+ # Recover cached decode-phase attention
+ # metadata structure
+ return self._cached_decode_metadata
+ assert ((self.seq_lens_tensor is not None)
+ or (self.encoder_seq_lens_tensor is not None))
+
+ # Compute some attn_metadata fields which default to None
+ slot_mapping = (None if self.slot_mapping is None else
+ self.slot_mapping[self.num_prefill_tokens:])
+ seq_lens_tensor = (None if self.seq_lens_tensor is None else
+ self.seq_lens_tensor[self.num_prefills:])
+ block_tables = (None if self.block_tables is None else
+ self.block_tables[self.num_prefills:])
+
+ # Construct & cache decode-phase attention metadata structure
+ self._cached_decode_metadata = XFormersMetadata(
+ num_prefills=0,
+ num_prefill_tokens=0,
+ num_decode_tokens=self.num_decode_tokens,
+ slot_mapping=slot_mapping,
+ seq_lens_tensor=seq_lens_tensor,
+ max_prefill_seq_len=0,
+ max_decode_seq_len=self.max_decode_seq_len,
+ block_tables=block_tables,
+ use_cuda_graph=self.use_cuda_graph,
+ # Begin encoder & cross attn fields below...
+ encoder_seq_lens=self.encoder_seq_lens,
+ encoder_seq_lens_tensor=self.encoder_seq_lens_tensor,
+ max_encoder_seq_len=self.max_encoder_seq_len,
+ cross_slot_mapping=self.cross_slot_mapping,
+ cross_block_tables=self.cross_block_tables)
+ return self._cached_decode_metadata
+
+
+def _get_attn_bias(
+ attn_metadata: XFormersMetadata,
+ attn_type: AttentionType,
+) -> Optional[AttentionBias]:
+ '''
+ Extract appropriate attention bias from attention metadata
+ according to attention type.
+
+ Arguments:
+
+ * attn_metadata: Attention metadata structure associated with attention
+ * attn_type: encoder attention, decoder self-attention,
+ encoder/decoder cross-attention
+
+ Returns:
+ * Appropriate attention bias value given the attention type
+ '''
+
+ if attn_type == AttentionType.DECODER:
+ return attn_metadata.attn_bias
+ elif attn_type == AttentionType.ENCODER:
+ return attn_metadata.encoder_attn_bias
+ else:
+ # attn_type == AttentionType.ENCODER_DECODER
+ return attn_metadata.cross_attn_bias
+
+
+def _set_attn_bias(
+ attn_metadata: XFormersMetadata,
+ attn_bias: List[Optional[AttentionBias]],
+ attn_type: AttentionType,
+) -> None:
+ '''
+ Update appropriate attention bias field of attention metadata,
+ according to attention type.
+
+ Arguments:
+
+ * attn_metadata: Attention metadata structure associated with attention
+ * attn_bias: The desired attention bias value
+ * attn_type: encoder attention, decoder self-attention,
+ encoder/decoder cross-attention
+ '''
+
+ if attn_type == AttentionType.DECODER:
+ attn_metadata.attn_bias = attn_bias
+ elif attn_type == AttentionType.ENCODER:
+ attn_metadata.encoder_attn_bias = attn_bias
+ elif attn_type == AttentionType.ENCODER_DECODER:
+ attn_metadata.cross_attn_bias = attn_bias
+ else:
+ raise AttributeError(f"Invalid attention type {str(attn_type)}")
+
+
+def _get_seq_len_block_table_args(
+ attn_metadata: XFormersMetadata,
+ is_prompt: bool,
+ attn_type: AttentionType,
+) -> tuple:
+ '''
+ The particular choice of sequence-length- and block-table-related
+ attributes which should be extracted from attn_metadata is dependent
+ on the type of attention operation.
+
+ Decoder attn -> select entirely decoder self-attention-related fields
+ Encoder/decoder cross-attn -> select encoder sequence lengths &
+ cross-attn block-tables fields
+ Encoder attn -> select encoder sequence lengths fields & no block tables
+
+ Arguments:
+
+ * attn_metadata: Attention metadata structure associated with attention op
+ * is_prompt: True if prefill, False otherwise
+ * attn_type: encoder attention, decoder self-attention,
+ encoder/decoder cross-attention
+
+ Returns:
+
+ * Appropriate sequence-lengths tensor
+ * Appropriate max sequence-length scalar
+ * Appropriate block tables (or None)
+ '''
+
+ if attn_type == AttentionType.DECODER:
+ # Decoder self-attention
+ # Choose max_seq_len based on whether we are in prompt_run
+ if is_prompt:
+ max_seq_len = attn_metadata.max_prefill_seq_len
+ else:
+ max_seq_len = attn_metadata.max_decode_seq_len
+ return (attn_metadata.seq_lens_tensor, max_seq_len,
+ attn_metadata.block_tables)
+ elif attn_type == AttentionType.ENCODER_DECODER:
+ # Enc/dec cross-attention KVs match encoder sequence length;
+ # cross-attention utilizes special "cross" block tables
+ return (attn_metadata.encoder_seq_lens_tensor,
+ attn_metadata.max_encoder_seq_len,
+ attn_metadata.cross_block_tables)
+ elif attn_type == AttentionType.ENCODER:
+ # No block tables associated with encoder attention
+ return (attn_metadata.encoder_seq_lens_tensor,
+ attn_metadata.max_encoder_seq_len, None)
+ else:
+ raise AttributeError(f"Invalid attention type {str(attn_type)}")
+
+
+class XFormersMetadataBuilder(CommonMetadataBuilder[XFormersMetadata]):
+
+ _metadata_cls = XFormersMetadata
+
+
+class XFormersImpl(AttentionImpl[XFormersMetadata]):
+ """
+ If the input tensors contain prompt tokens, the layout is as follows:
+ |<--------------- num_prefill_tokens ----------------->|
+ |<--prefill_0-->|<--prefill_1-->|...|<--prefill_N-1--->|
+
+ Otherwise, the layout is as follows:
+ |<----------------- num_decode_tokens ------------------>|
+ |<--decode_0-->|..........|<--decode_M-1-->|<--padding-->|
+
+ Generation tokens can contain padding when cuda-graph is used.
+ Currently, prompt tokens don't contain any padding.
+
+ The prompts might have different lengths, while the generation tokens
+ always have length 1.
+
+ If chunked prefill is enabled, prefill tokens and decode tokens can be
+ batched together in a flattened 1D query.
+
+ |<----- num_prefill_tokens ---->|<------- num_decode_tokens --------->|
+ |<-prefill_0->|...|<-prefill_N-1->|<--decode_0-->|...|<--decode_M-1-->|
+
+ Currently, cuda graph is disabled for chunked prefill, meaning there's no
+ padding between prefill and decode tokens.
+ """
+
+ def __init__(
+ self,
+ num_heads: int,
+ head_size: int,
+ scale: float,
+ num_kv_heads: int,
+ alibi_slopes: Optional[List[float]],
+ sliding_window: Optional[int],
+ kv_cache_dtype: str,
+ blocksparse_params: Optional[Dict[str, Any]] = None,
+ logits_soft_cap: Optional[float] = None,
+ ) -> None:
+ if blocksparse_params is not None:
+ raise ValueError(
+ "XFormers does not support block-sparse attention.")
+ if logits_soft_cap is not None:
+ raise ValueError(
+ "XFormers does not support attention logits soft capping.")
+ self.num_heads = num_heads
+ self.head_size = head_size
+ self.scale = float(scale)
+ self.num_kv_heads = num_kv_heads
+ if alibi_slopes is not None:
+ alibi_slopes = torch.tensor(alibi_slopes, dtype=torch.float32)
+ self.alibi_slopes = alibi_slopes
+ self.sliding_window = sliding_window
+ self.kv_cache_dtype = kv_cache_dtype
+
+ assert self.num_heads % self.num_kv_heads == 0
+ self.num_queries_per_kv = self.num_heads // self.num_kv_heads
+
+ suppored_head_sizes = PagedAttention.get_supported_head_sizes()
+ if head_size not in suppored_head_sizes:
+ raise ValueError(
+ f"Head size {head_size} is not supported by PagedAttention. "
+ f"Supported head sizes are: {suppored_head_sizes}.")
+ self.head_mapping = torch.repeat_interleave(
+ torch.arange(self.num_kv_heads, dtype=torch.int32),
+ self.num_queries_per_kv)
+
+ def forward(
+ self,
+ query: torch.Tensor,
+ key: Optional[torch.Tensor],
+ value: Optional[torch.Tensor],
+ kv_cache: torch.Tensor,
+ attn_metadata: "XFormersMetadata",
+ k_scale: float = 1.0,
+ v_scale: float = 1.0,
+ attn_type: AttentionType = AttentionType.DECODER,
+ ) -> torch.Tensor:
+ """Forward pass with xFormers and PagedAttention.
+
+ For decoder-only models: query, key and value must be non-None.
+
+ For encoder/decoder models:
+ * XFormersImpl.forward() may be invoked for both self- and cross-
+ attention layers.
+ * For self-attention: query, key and value must be non-None.
+ * For cross-attention:
+ * Query must be non-None
+ * During prefill, key and value must be non-None; key and value
+ get cached for use during decode.
+ * During decode, key and value may be None, since:
+ (1) key and value tensors were cached during prefill, and
+ (2) cross-attention key and value tensors do not grow during
+ decode
+
+ A note on how the attn_type (attention type enum) argument impacts
+ attention forward() behavior:
+
+ * DECODER: normal decoder-only behavior;
+ use decoder self-attention block table
+ * ENCODER: no KV caching; pass encoder sequence
+ attributes (encoder_seq_lens/encoder_seq_lens_tensor/
+ max_encoder_seq_len) to kernel, in lieu of decoder
+ sequence attributes (seq_lens/seq_lens_tensor/max_seq_len)
+ * ENCODER_DECODER: cross-attention behavior;
+ use cross-attention block table for caching KVs derived
+ from encoder hidden states; since KV sequence lengths
+ will match encoder sequence lengths, pass encoder sequence
+ attributes to kernel (encoder_seq_lens/encoder_seq_lens_tensor/
+ max_encoder_seq_len)
+
+ Args:
+ query: shape = [num_tokens, num_heads * head_size]
+ key: shape = [num_tokens, num_kv_heads * head_size]
+ value: shape = [num_tokens, num_kv_heads * head_size]
+ kv_cache = [2, num_blocks, block_size * num_kv_heads * head_size]
+ NOTE: kv_cache will be an empty tensor with shape [0]
+ for profiling run.
+ attn_metadata: Metadata for attention.
+ attn_type: Select attention type, between encoder attention,
+ decoder self-attention, or encoder/decoder cross-
+ attention. Defaults to decoder self-attention,
+ which is the vLLM default generally
+ Returns:
+ shape = [num_tokens, num_heads * head_size]
+ """
+
+ # Check that appropriate attention metadata attributes are
+ # selected for the desired attention type
+ if (attn_type == AttentionType.ENCODER
+ and (not attn_metadata.is_all_encoder_attn_metadata_set)):
+ raise AttributeError("Encoder attention requires setting "
+ "encoder metadata attributes.")
+ elif (attn_type == AttentionType.ENCODER_DECODER
+ and (not attn_metadata.is_all_cross_attn_metadata_set)):
+ raise AttributeError("Encoder/decoder cross-attention "
+ "requires setting cross-attention "
+ "metadata attributes.")
+
+ query = query.view(-1, self.num_heads, self.head_size)
+ if key is not None:
+ assert value is not None
+ key = key.view(-1, self.num_kv_heads, self.head_size)
+ value = value.view(-1, self.num_kv_heads, self.head_size)
+ else:
+ assert value is None
+
+ # Self-attention vs. cross-attention will impact
+ # which KV cache memory-mapping & which
+ # seqlen datastructures we utilize
+
+ if (attn_type != AttentionType.ENCODER and kv_cache.numel() > 0):
+ # KV-cache during decoder-self- or
+ # encoder-decoder-cross-attention, but not
+ # during encoder attention.
+ #
+ # Even if there are no new key/value pairs to cache,
+ # we still need to break out key_cache and value_cache
+ # i.e. for later use by paged attention
+ key_cache, value_cache = PagedAttention.split_kv_cache(
+ kv_cache, self.num_kv_heads, self.head_size)
+
+ if (key is not None) and (value is not None):
+
+ if attn_type == AttentionType.ENCODER_DECODER:
+ # Update cross-attention KV cache (prefill-only)
+ # During cross-attention decode, key & value will be None,
+ # preventing this IF-statement branch from running
+ updated_slot_mapping = attn_metadata.cross_slot_mapping
+ else:
+ # Update self-attention KV cache (prefill/decode)
+ updated_slot_mapping = attn_metadata.slot_mapping
+
+ # Reshape the input keys and values and store them in the cache.
+ # If kv_cache is not provided, the new key and value tensors are
+ # not cached. This happens during the initial memory
+ # profiling run.
+ PagedAttention.write_to_paged_cache(key, value, key_cache,
+ value_cache,
+ updated_slot_mapping,
+ self.kv_cache_dtype,
+ k_scale, v_scale)
+
+ if attn_type == AttentionType.ENCODER:
+ # Encoder attention - chunked prefill is not applicable;
+ # derive token-count from query shape & and treat them
+ # as 100% prefill tokens
+ assert attn_metadata.num_encoder_tokens is not None
+ num_prefill_tokens = attn_metadata.num_encoder_tokens
+ num_encoder_tokens = attn_metadata.num_encoder_tokens
+ num_decode_tokens = 0
+ elif attn_type == AttentionType.DECODER:
+ # Decoder self-attention supports chunked prefill.
+ num_prefill_tokens = attn_metadata.num_prefill_tokens
+ num_encoder_tokens = attn_metadata.num_prefill_tokens
+ num_decode_tokens = attn_metadata.num_decode_tokens
+ # Only enforce this shape-constraint for decoder
+ # self-attention
+ assert key.shape[0] == num_prefill_tokens + num_decode_tokens
+ assert value.shape[0] == num_prefill_tokens + num_decode_tokens
+ else: # attn_type == AttentionType.ENCODER_DECODER
+ # Encoder/decoder cross-attention requires no chunked
+ # prefill (100% prefill or 100% decode tokens, no mix)
+ num_prefill_tokens = attn_metadata.num_prefill_tokens
+ if attn_metadata.num_encoder_tokens is not None:
+ num_encoder_tokens = attn_metadata.num_encoder_tokens
+ else:
+ num_encoder_tokens = attn_metadata.num_prefill_tokens
+ num_decode_tokens = attn_metadata.num_decode_tokens
+ output = torch.empty_like(query)
+ # Query for decode. KV is not needed because it is already cached.
+ decode_query = query[num_prefill_tokens:]
+ # QKV for prefill.
+ query = query[:num_prefill_tokens]
+ if key is not None and value is not None:
+ key = key[:num_encoder_tokens]
+ value = value[:num_encoder_tokens]
+ assert query.shape[0] == num_prefill_tokens
+ assert decode_query.shape[0] == num_decode_tokens
+
+ if prefill_meta := attn_metadata.prefill_metadata:
+ # Prompt run.
+ if kv_cache.numel() == 0 or prefill_meta.block_tables.numel() == 0:
+ # normal attention.
+ # block tables are empty if the prompt does not have a cached
+ # prefix.
+ out = self._run_memory_efficient_xformers_forward(
+ query, key, value, prefill_meta, attn_type=attn_type)
+ assert out.shape == output[:num_prefill_tokens].shape
+ output[:num_prefill_tokens] = out
+ else:
+
+ assert prefill_meta.query_start_loc is not None
+ assert prefill_meta.max_query_len is not None
+
+ # prefix-enabled attention
+ # TODO(Hai) this triton kernel has regression issue (broke) to
+ # deal with different data types between KV and FP8 KV cache,
+ # to be addressed separately.
+ out = PagedAttention.forward_prefix(
+ query,
+ key,
+ value,
+ self.kv_cache_dtype,
+ key_cache,
+ value_cache,
+ prefill_meta.block_tables,
+ prefill_meta.query_start_loc,
+ prefill_meta.seq_lens_tensor,
+ prefill_meta.context_lens_tensor,
+ prefill_meta.max_query_len,
+ self.alibi_slopes,
+ self.sliding_window,
+ k_scale,
+ v_scale,
+ )
+ assert output[:num_prefill_tokens].shape == out.shape
+ output[:num_prefill_tokens] = out
+
+ if decode_meta := attn_metadata.decode_metadata:
+
+ (
+ seq_lens_arg,
+ max_seq_len_arg,
+ block_tables_arg,
+ ) = _get_seq_len_block_table_args(decode_meta, False, attn_type)
+
+ output[num_prefill_tokens:] = PagedAttention.forward_decode(
+ decode_query,
+ key_cache,
+ value_cache,
+ block_tables_arg,
+ seq_lens_arg,
+ max_seq_len_arg,
+ self.kv_cache_dtype,
+ self.head_mapping,
+ self.scale,
+ self.alibi_slopes,
+ k_scale,
+ v_scale,
+ )
+
+ # Reshape the output tensor.
+ return output.view(-1, self.num_heads * self.head_size)
+
+ def _run_sdpa_fallback(
+ self,
+ query: torch.Tensor,
+ key: torch.Tensor,
+ value: torch.Tensor,
+ attn_metadata: "XFormersMetadata",
+ ) -> torch.Tensor:
+ """Pure-math causal attention fallback with Q-tiling memory optimization.
+
+ Called when: kv_cache.numel()==0 (profiling) AND head_size > 128.
+ No KV cache prefix in this path — KV length == query length.
+
+ Memory optimization (Q-tiling, same principle as Flash Attention):
+ Split Q into _Q_CHUNK-sized sub-blocks, compute per-block,
+ peak memory O(_Q_CHUNK × q_len) instead of O(q_len²).
+
+ Softmax computed in float32 to prevent float16 overflow.
+
+ Args:
+ query : [1, total_query_tokens, num_heads, head_dim]
+ key : [1, total_query_tokens, num_kv_heads, head_dim]
+ value : [1, total_query_tokens, num_kv_heads, head_dim]
+ Returns:
+ [1, total_query_tokens, num_heads, head_dim]
+ """
+ _Q_CHUNK = 256
+
+ assert attn_metadata.seq_lens is not None
+ orig_dtype = query.dtype
+ num_seqs = len(attn_metadata.seq_lens)
+
+ if (attn_metadata.query_start_loc is not None
+ and len(attn_metadata.query_start_loc) == num_seqs + 1):
+ q_lens = [
+ int(attn_metadata.query_start_loc[i + 1].item()) -
+ int(attn_metadata.query_start_loc[i].item())
+ for i in range(num_seqs)
+ ]
+ else:
+ q_lens = list(attn_metadata.seq_lens)
+
+ q_flat = query.squeeze(0)
+ k_flat = key.squeeze(0)
+ v_flat = value.squeeze(0)
+
+ output = torch.empty_like(q_flat)
+ seq_start = 0
+ for q_len in q_lens:
+ seq_end = seq_start + q_len
+
+ 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()
+
+ if k_s.shape[0] != self.num_heads:
+ n = self.num_heads // k_s.shape[0]
+ k_s = k_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)
+
+ for qc_start in range(0, q_len, _Q_CHUNK):
+ qc_end = min(qc_start + _Q_CHUNK, q_len)
+
+ q_c = q_flat[seq_start + qc_start:seq_start + qc_end] \
+ .permute(1, 0, 2).float()
+
+ 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)
+ mask = k_pos.unsqueeze(0) > qc_q_pos.unsqueeze(1)
+ attn_w = attn_w.masked_fill(mask.unsqueeze(0), float("-inf"))
+
+ attn_w = torch.softmax(attn_w, dim=-1)
+ out_c = torch.matmul(attn_w, v_s).to(orig_dtype)
+
+ output[seq_start + qc_start:seq_start + qc_end] = (
+ out_c.permute(1, 0, 2))
+
+ seq_start = seq_end
+
+ return output.unsqueeze(0)
+
+ def _run_memory_efficient_xformers_forward(
+ self,
+ query: torch.Tensor,
+ key: torch.Tensor,
+ value: torch.Tensor,
+ attn_metadata: XFormersMetadata,
+ attn_type: AttentionType = AttentionType.DECODER,
+ ) -> torch.Tensor:
+ """Attention for 1D query of multiple prompts. Multiple prompt
+ tokens are flattened in to `query` input.
+
+ See https://facebookresearch.github.io/xformers/components/ops.html
+ for API spec.
+
+ Args:
+ output: shape = [num_prefill_tokens, num_heads, head_size]
+ query: shape = [num_prefill_tokens, num_heads, head_size]
+ key: shape = [num_prefill_tokens, num_kv_heads, head_size]
+ value: shape = [num_prefill_tokens, num_kv_heads, head_size]
+ attn_metadata: Metadata for attention.
+ attn_type: Select attention type, between encoder attention,
+ decoder self-attention, or encoder/decoder cross-
+ attention. Defaults to decoder self-attention,
+ which is the vLLM default generally
+ """
+
+ original_query = query
+ # if self.num_kv_heads != self.num_heads:
+ # # GQA/MQA requires the shape [B, M, G, H, K].
+ # # Note that the output also has the same shape (which is different
+ # # from a spec from the doc).
+ # query = query.view(query.shape[0], self.num_kv_heads,
+ # self.num_queries_per_kv, query.shape[-1])
+ # print(f"5555555555555 q shape {query.shape}")
+ # key = key[:, :,
+ # None, :].expand(key.shape[0], self.num_kv_heads,
+ # self.num_queries_per_kv, key.shape[-1])
+ # value = value[:, :,
+ # None, :].expand(value.shape[0], self.num_kv_heads,
+ # self.num_queries_per_kv,
+ # value.shape[-1])
+ # Set attention bias if not provided. This typically happens at
+ # the very attention layer of every iteration.
+ # FIXME(woosuk): This is a hack.
+ attn_bias = _get_attn_bias(attn_metadata, attn_type)
+ if attn_bias is None:
+ if self.alibi_slopes is None:
+ if (attn_type == AttentionType.ENCODER_DECODER):
+ assert attn_metadata.seq_lens is not None
+ assert attn_metadata.encoder_seq_lens is not None
+
+ # Default enc/dec cross-attention mask is non-causal
+ attn_bias = BlockDiagonalMask.from_seqlens(
+ attn_metadata.seq_lens, attn_metadata.encoder_seq_lens)
+ elif attn_type == AttentionType.ENCODER:
+ assert attn_metadata.encoder_seq_lens is not None
+
+ # Default encoder self-attention mask is non-causal
+ attn_bias = BlockDiagonalMask.from_seqlens(
+ attn_metadata.encoder_seq_lens)
+ else:
+ assert attn_metadata.seq_lens is not None
+
+ # Default decoder self-attention mask is causal
+ attn_bias = BlockDiagonalCausalMask.from_seqlens(
+ attn_metadata.seq_lens)
+ if self.sliding_window is not None:
+ attn_bias = attn_bias.make_local_attention(
+ self.sliding_window)
+ attn_bias = [attn_bias]
+ else:
+ assert attn_metadata.seq_lens is not None
+ attn_bias = _make_alibi_bias(self.alibi_slopes,
+ self.num_kv_heads, query.dtype,
+ attn_metadata.seq_lens)
+
+ _set_attn_bias(attn_metadata, attn_bias, attn_type)
+
+ # No alibi slopes.
+ # TODO(woosuk): Too many view operations. Let's try to reduce
+ # them in the future for code readability.
+ self.attn_op = xops.fmha.flash.FwOp()
+ if self.alibi_slopes is None:
+ # Add the batch dimension.
+ query = query.unsqueeze(0)
+ key = key.unsqueeze(0)
+ value = value.unsqueeze(0)
+ if self.head_size > 128:
+ out = self._run_sdpa_fallback(query, key, value, attn_metadata)
+ else:
+ out = xops.memory_efficient_attention_forward(
+ query,
+ key,
+ value,
+ attn_bias=attn_bias[0],
+ p=0.0,
+ scale=self.scale,
+ op=self.attn_op,
+ )
+ return out.view_as(original_query)
+
+ # Attention with alibi slopes.
+ # FIXME(woosuk): Because xformers does not support dynamic sequence
+ # lengths with custom attention bias, we process each prompt one by
+ # one. This is inefficient, especially when we have many short prompts.
+ assert attn_metadata.seq_lens is not None
+ output = torch.empty_like(original_query)
+ start = 0
+ for i, seq_len in enumerate(attn_metadata.seq_lens):
+ end = start + seq_len
+ out = xops.memory_efficient_attention_forward(
+ query[None, start:end],
+ key[None, start:end],
+ value[None, start:end],
+ attn_bias=attn_bias[i],
+ p=0.0,
+ scale=self.scale,
+ )
+ # TODO(woosuk): Unnecessary copy. Optimize.
+ output[start:end].copy_(out.view_as(original_query[start:end]))
+ start += seq_len
+ return output
+
+
+def _make_alibi_bias(
+ alibi_slopes: torch.Tensor,
+ num_kv_heads: int,
+ dtype: torch.dtype,
+ seq_lens: List[int],
+) -> List[AttentionBias]:
+ attn_biases: List[AttentionBias] = []
+ for seq_len in seq_lens:
+ bias = torch.arange(seq_len, dtype=dtype)
+ # NOTE(zhuohan): HF uses
+ # `bias = bias[None, :].repeat(seq_len, 1)`
+ # here. We find that both biases give the same results, but
+ # the bias below more accurately follows the original ALiBi
+ # paper.
+ # Calculate a matrix where each element represents ith element- jth
+ # element.
+ bias = bias[None, :] - bias[:, None]
+
+ padded_len = (seq_len + 7) // 8 * 8
+ num_heads = alibi_slopes.shape[0]
+ bias = torch.empty(
+ 1, # batch size
+ num_heads,
+ seq_len,
+ padded_len,
+ device=alibi_slopes.device,
+ dtype=dtype,
+ )[:, :, :, :seq_len].copy_(bias)
+ bias.mul_(alibi_slopes[:, None, None])
+ if num_heads != num_kv_heads:
+ bias = bias.unflatten(1, (num_kv_heads, num_heads // num_kv_heads))
+ attn_biases.append(LowerTriangularMaskWithTensorBias(bias))
+
+ return attn_biases
\ No newline at end of file