clef / code /models /common /llm_runtime /vllm_adapter.py
tt-hous's picture
Add files using upload-large-folder tool
2415c4c verified
Raw History Blame Contribute Delete
17.1 kB
# SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
# SPDX-License-Identifier: Apache-2.0
"""Normalization for the external vLLM call and KV-cache contracts."""
from __future__ import annotations
import dataclasses
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Any, TypedDict
import torch
from models.common.llm_runtime.config import PagedKVCacheConfig, TraceConfig
from models.common.llm_runtime.paged_kv_cache import torch_dtype_for_ttnn
# These plugin fields are meaningful to other model families, but the registered
# text-only TTTv2 paths do not implement their hybrid-cache or mRoPE semantics.
# Sampling history and slot-remap fields are intentionally *not* ignored: they
# carry request-owned penalty and RNG lifecycle state through the common runtime.
_IGNORED_VLLM_KWARGS = frozenset(
{
"page_tables_per_layer",
"rope_deltas_all_users",
}
)
_SAMPLING_STATE_VLLM_FIELD_ORDER = ("prompt_tokens", "output_tokens", "slot_remap")
_SAMPLING_STATE_VLLM_KWARGS = frozenset(_SAMPLING_STATE_VLLM_FIELD_ORDER)
class _NormalizedPrefillRequiredKwargs(TypedDict):
tokens: torch.Tensor # ↓ Core request
page_table: torch.Tensor
prompt_lens: torch.Tensor | None # ↓ Sequence metadata
start_pos: torch.Tensor | None
empty_slots: Sequence[int] | None # ↓ Lane routing
kv_cache: Any # ↓ Borrowed resources
sampling_params: Any # ↓ Sampling
class NormalizedPrefillKwargs(_NormalizedPrefillRequiredKwargs, total=False):
prompt_tokens: Any # ↓ Request-owned sampling state
output_tokens: Any
slot_remap: Any
class _NormalizedDecodeRequiredKwargs(TypedDict):
tokens: torch.Tensor # ↓ Core request
start_pos: torch.Tensor
page_table: torch.Tensor
kv_cache: Any # ↓ Borrowed resources
sampling_params: Any # ↓ Sampling
reset_batch: bool # ↓ State transition
class NormalizedDecodeKwargs(_NormalizedDecodeRequiredKwargs, total=False):
prompt_tokens: Any # ↓ Request-owned sampling state
output_tokens: Any
slot_remap: Any
@dataclass(frozen=True)
class VLLMAdapterConfig:
"""Fully resolved static policy for the vLLM compatibility boundary."""
trace: TraceConfig
paged_kv_cache: PagedKVCacheConfig
expected_num_layers: int
expected_kv_heads_per_device: int | None
expected_head_dim: int | None
model_kv_cache_dtypes: tuple[Any, ...]
request_state_fields: tuple[str, ...] = ()
def __post_init__(self) -> None:
_validate_resolved_adapter_config(self)
@classmethod
def resolve(
cls,
*,
trace: TraceConfig,
paged_kv_cache: PagedKVCacheConfig,
expected_num_layers: int,
expected_kv_heads_per_device: int | None = None,
expected_head_dim: int | None = None,
model_kv_cache_dtype: Any | Sequence[Any],
request_state_fields: Sequence[str] = (),
) -> "VLLMAdapterConfig":
if type(trace) is not TraceConfig:
raise TypeError("trace must be a TraceConfig")
if type(paged_kv_cache) is not PagedKVCacheConfig:
raise TypeError("paged_kv_cache must be a PagedKVCacheConfig")
resolved_num_layers = _resolve_positive_int("expected_num_layers", expected_num_layers)
resolved_kv_heads = _resolve_optional_positive_int("expected_kv_heads_per_device", expected_kv_heads_per_device)
resolved_head_dim = _resolve_optional_positive_int("expected_head_dim", expected_head_dim)
if model_kv_cache_dtype is None:
raise TypeError("model_kv_cache_dtype must be supplied from model metadata")
if isinstance(model_kv_cache_dtype, Sequence) and not isinstance(model_kv_cache_dtype, (str, bytes)):
model_kv_cache_dtypes = tuple(model_kv_cache_dtype)
else:
model_kv_cache_dtypes = (model_kv_cache_dtype,)
return cls(
trace=trace,
paged_kv_cache=paged_kv_cache,
expected_num_layers=resolved_num_layers,
expected_kv_heads_per_device=resolved_kv_heads,
expected_head_dim=resolved_head_dim,
model_kv_cache_dtypes=model_kv_cache_dtypes,
request_state_fields=_resolve_request_state_fields(request_state_fields),
)
class VLLMAdapter:
"""Convert vLLM-facing calls into the configured runtime call surface.
`Llama3Generator` calls `normalize_prefill` or
`normalize_decode` before choosing an eager or traced execution
target. Each call returns typed host tensors and its validated Boolean
trace choice separately. During vLLM initialization,
`resolve_legacy_kv_cache_config` validates the external cache shape
and returns the one allowed capacity-resolved runtime configuration.
The adapter owns no TT tensors or model/runtime resources. Model-specific
construction supplies the already-derived KV shape and dtype expectations.
"""
def __init__(self, config: VLLMAdapterConfig) -> None:
if not isinstance(config, VLLMAdapterConfig):
raise TypeError("config must be a VLLMAdapterConfig")
self.config = config
# Public API
def normalize_prefill(
self,
tokens: torch.Tensor,
page_table: torch.Tensor,
*,
enable_trace: bool, # ↓ Required policy
prompt_lens: Sequence[int] | torch.Tensor | None = None, # ↓ Sequence metadata
start_pos: torch.Tensor | None = None,
empty_slots: Sequence[int] | None = None, # ↓ Lane routing
kv_cache: Any = None, # ↓ Borrowed resources
sampling_params: Any = None, # ↓ Sampling
compatibility_kwargs: Mapping[str, Any] | None = None, # ↓ Compatibility
) -> tuple[NormalizedPrefillKwargs, bool]:
"""Normalize one explicit prefill request and return trace intent separately."""
self._validate_compatibility_kwargs(compatibility_kwargs, operation="prefill")
self._validate_trace_selection(enable_trace, operation="prefill")
normalized: NormalizedPrefillKwargs = {
"tokens": tokens,
"page_table": page_table,
"prompt_lens": prompt_lens,
"start_pos": start_pos,
"empty_slots": empty_slots,
"kv_cache": kv_cache,
"sampling_params": sampling_params,
}
_copy_supplied_sampling_state(normalized, compatibility_kwargs, self.config.request_state_fields)
_normalize_tensor(normalized, "tokens", torch.long)
_normalize_tensor(normalized, "page_table", torch.int32)
_normalize_tensor(normalized, "prompt_lens", torch.long)
_normalize_tensor(normalized, "start_pos", torch.long)
return normalized, enable_trace
def normalize_decode(
self,
tokens: torch.Tensor,
start_pos: torch.Tensor,
page_table: torch.Tensor,
*,
enable_trace: bool, # ↓ Required policy
kv_cache: Any = None, # ↓ Borrowed resources
sampling_params: Any = None, # ↓ Sampling
reset_batch: bool = False, # ↓ State transition
compatibility_kwargs: Mapping[str, Any] | None = None, # ↓ Compatibility
) -> tuple[NormalizedDecodeKwargs, bool]:
"""Normalize one explicit decode request and return trace intent separately."""
self._validate_compatibility_kwargs(compatibility_kwargs, operation="decode")
self._validate_trace_selection(enable_trace, operation="decode")
normalized: NormalizedDecodeKwargs = {
"tokens": tokens,
"start_pos": start_pos,
"page_table": page_table,
"kv_cache": kv_cache,
"sampling_params": sampling_params,
"reset_batch": reset_batch,
}
_copy_supplied_sampling_state(normalized, compatibility_kwargs, self.config.request_state_fields)
_normalize_tensor(normalized, "tokens", torch.long)
_normalize_tensor(normalized, "start_pos", torch.long)
_normalize_tensor(normalized, "page_table", torch.int32)
tokens = normalized["tokens"]
if tokens.ndim == 2 and tokens.shape[-1] == 1:
normalized["tokens"] = tokens.reshape(-1)
return normalized, enable_trace
def resolve_legacy_kv_cache_config(
self,
kv_cache_shape: Sequence[int],
dtype: torch.dtype,
num_layers: int,
) -> PagedKVCacheConfig:
"""Validate vLLM's legacy KV spec and return a resolved frozen config."""
shape = tuple(int(dim) for dim in kv_cache_shape)
if len(shape) != 4:
raise ValueError(f"KV cache shape must have rank 4, got {shape}")
num_blocks, kv_heads, block_size, head_dim = shape
config = self.config.paged_kv_cache
if num_blocks <= 0:
raise ValueError("KV cache num_blocks must be positive")
if block_size <= 0:
raise ValueError("KV cache block_size must be positive")
if (
self.config.expected_kv_heads_per_device is not None
and kv_heads != self.config.expected_kv_heads_per_device
):
raise ValueError(
f"vLLM KV heads {kv_heads} do not match model-derived KV heads "
f"{self.config.expected_kv_heads_per_device}"
)
if self.config.expected_head_dim is not None and head_dim != self.config.expected_head_dim:
raise ValueError(
f"vLLM KV head dimension {head_dim} does not match model-derived head dimension "
f"{self.config.expected_head_dim}"
)
if int(num_layers) != self.config.expected_num_layers:
raise ValueError(
f"vLLM KV layer count {num_layers} does not match model-derived layer count "
f"{self.config.expected_num_layers}"
)
_validate_vllm_torch_dtype(dtype, self.config.model_kv_cache_dtypes)
configured_num_blocks = config.num_blocks
if configured_num_blocks is not None and int(configured_num_blocks) != num_blocks:
raise ValueError(
f"PagedKVCacheConfig is already resolved to {configured_num_blocks} blocks; "
f"vLLM requested {num_blocks}"
)
return dataclasses.replace(
config,
block_size=block_size,
max_num_blocks=num_blocks,
num_blocks=num_blocks,
)
# Private implementation
@staticmethod
def _validate_compatibility_kwargs(
compatibility_kwargs: Mapping[str, Any] | None,
*,
operation: str,
) -> None:
if compatibility_kwargs is None:
return
if not isinstance(compatibility_kwargs, Mapping):
raise TypeError("compatibility_kwargs must be a mapping")
supported_keys = _IGNORED_VLLM_KWARGS | _SAMPLING_STATE_VLLM_KWARGS
unknown_keys = sorted(key for key in compatibility_kwargs if key not in supported_keys)
if unknown_keys:
raise TypeError(f"{operation} got an unexpected keyword argument {unknown_keys[0]!r}")
def _validate_trace_selection(self, enable_trace: bool, *, operation: str) -> None:
if not isinstance(enable_trace, bool):
raise TypeError("enable_trace must be bool")
if operation == "prefill":
configured = self.config.trace.prefill_enabled
elif operation == "decode":
configured = self.config.trace.decode_enabled
else:
raise ValueError(f"Unknown trace operation {operation!r}")
# TraceConfig describes the trace targets made available at construction.
# The unchanged vLLM loader constructs the default ``all`` superset and
# presents its effective static mode through these required per-operation
# booleans. Reject only requests for a trace target that was not built;
# an unselected operation intentionally uses eager execution.
if enable_trace and not configured:
raise ValueError(
f"enable_trace={enable_trace} for {operation} disagrees with static "
f"TraceConfig policy ({configured})"
)
def _resolve_positive_int(name: str, value: Any) -> int:
if isinstance(value, bool):
raise ValueError(f"{name} must be a positive integer")
try:
resolved = int(value)
except (TypeError, ValueError) as error:
raise ValueError(f"{name} must be a positive integer") from error
if resolved <= 0 or resolved != value:
raise ValueError(f"{name} must be a positive integer")
return resolved
def _resolve_optional_positive_int(name: str, value: Any) -> int | None:
return None if value is None else _resolve_positive_int(name, value)
def _validate_resolved_adapter_config(config: VLLMAdapterConfig) -> None:
if type(config.trace) is not TraceConfig:
raise TypeError("trace must be a TraceConfig")
if type(config.paged_kv_cache) is not PagedKVCacheConfig:
raise TypeError("paged_kv_cache must be a PagedKVCacheConfig")
_require_resolved_positive_int("expected_num_layers", config.expected_num_layers)
if config.expected_kv_heads_per_device is not None:
_require_resolved_positive_int("expected_kv_heads_per_device", config.expected_kv_heads_per_device)
if config.expected_head_dim is not None:
_require_resolved_positive_int("expected_head_dim", config.expected_head_dim)
if not isinstance(config.model_kv_cache_dtypes, tuple):
raise TypeError("model_kv_cache_dtypes must be a tuple")
if not config.model_kv_cache_dtypes:
raise ValueError("model_kv_cache_dtypes cannot be empty")
if config.request_state_fields != _resolve_request_state_fields(config.request_state_fields):
raise ValueError("request_state_fields must be canonical")
if len(config.model_kv_cache_dtypes) not in (1, config.expected_num_layers):
raise ValueError("model_kv_cache_dtypes must be uniform or contain one dtype per model layer")
first_dtype = config.model_kv_cache_dtypes[0]
uniform_model_dtype = (
first_dtype if all(dtype == first_dtype for dtype in config.model_kv_cache_dtypes[1:]) else None
)
if uniform_model_dtype is not None and config.paged_kv_cache.dtype != uniform_model_dtype:
raise ValueError(
"PagedKVCacheConfig.dtype does not match the model-owned KV cache dtype: "
f"{config.paged_kv_cache.dtype!r} != {uniform_model_dtype!r}"
)
def _require_resolved_positive_int(name: str, value: Any) -> None:
if not isinstance(value, int) or isinstance(value, bool) or value <= 0:
raise ValueError(f"{name} must be a positive integer")
def _copy_supplied_sampling_state(
normalized: NormalizedPrefillKwargs | NormalizedDecodeKwargs,
compatibility_kwargs: Mapping[str, Any] | None,
request_state_fields: Sequence[str],
) -> None:
"""Forward only request state that the external caller actually supplied."""
if compatibility_kwargs is None:
return
for name in request_state_fields:
value = compatibility_kwargs.get(name)
if value is not None:
normalized[name] = value
def _resolve_request_state_fields(request_state_fields: Sequence[str]) -> tuple[str, ...]:
if isinstance(request_state_fields, (str, bytes)) or not isinstance(request_state_fields, Sequence):
raise TypeError("request_state_fields must be a sequence of field names")
resolved = tuple(request_state_fields)
if len(set(resolved)) != len(resolved):
raise ValueError("request_state_fields must not contain duplicates")
if any(name not in _SAMPLING_STATE_VLLM_KWARGS for name in resolved):
raise ValueError(f"request_state_fields must be drawn from {_SAMPLING_STATE_VLLM_FIELD_ORDER}")
if resolved != tuple(name for name in _SAMPLING_STATE_VLLM_FIELD_ORDER if name in resolved):
raise ValueError(f"request_state_fields must retain canonical order {_SAMPLING_STATE_VLLM_FIELD_ORDER}")
return resolved
def _normalize_tensor(kwargs: dict[str, Any], name: str, dtype: torch.dtype) -> None:
value = kwargs.get(name)
if value is None:
return
if isinstance(value, torch.Tensor):
kwargs[name] = value.to(dtype=dtype)
else:
kwargs[name] = torch.as_tensor(value, dtype=dtype)
def _validate_vllm_torch_dtype(dtype: torch.dtype, model_dtypes: tuple[Any, ...]) -> None:
if not isinstance(dtype, torch.dtype):
raise TypeError(f"vLLM KV dtype must be a torch.dtype, got {type(dtype).__name__}")
expected = {
model_dtype if isinstance(model_dtype, torch.dtype) else torch_dtype_for_ttnn(model_dtype)
for model_dtype in model_dtypes
}
if dtype not in expected or len(expected) != 1:
expected_names = ", ".join(sorted(str(item) for item in expected))
raise ValueError(f"vLLM KV dtype {dtype} does not match model-owned dtype surrogate {expected_names}")