"""PyTorch/Transformers adapter for tensor-verified KVC1 containers.""" from __future__ import annotations from pathlib import Path from typing import Any, Mapping import torch from transformers import DynamicCache import kvcache TORCH_TO_KVC = { torch.float16: "f16", torch.bfloat16: "bf16", torch.float32: "f32", torch.int8: "i8", torch.uint8: "u8", } KVC_TO_TORCH = {value: key for key, value in TORCH_TO_KVC.items()} def _cache_layers(cache: Any) -> list[tuple[torch.Tensor, torch.Tensor]]: raw_layers = getattr(cache, "layers", None) if not isinstance(raw_layers, (list, tuple)) or not raw_layers: raise kvcache.CacheFormatError("dynamic cache must expose at least one layer") layers: list[tuple[torch.Tensor, torch.Tensor]] = [] for index, layer in enumerate(raw_layers): keys = getattr(layer, "keys", None) values = getattr(layer, "values", None) if not isinstance(keys, torch.Tensor) or not isinstance(values, torch.Tensor): raise kvcache.CacheFormatError(f"dynamic cache layer {index} must expose key and value tensors") if keys.shape != values.shape: raise kvcache.CacheFormatError(f"dynamic cache layer {index} key and value tensors must have the same shape") if keys.dtype != values.dtype: raise kvcache.CacheFormatError(f"dynamic cache layer {index} key and value tensors must have the same dtype") if keys.ndim != 4: raise kvcache.CacheFormatError( f"dynamic cache layer {index} tensors must have [batch, heads, sequence, head_dim] shape" ) layers.append((keys.detach().cpu().contiguous(), values.detach().cpu().contiguous())) first_shape = layers[0][0].shape first_dtype = layers[0][0].dtype for index, (keys, _) in enumerate(layers[1:], start=1): if keys.shape != first_shape: raise kvcache.CacheFormatError(f"dynamic cache layer {index} tensors must have the same shape as layer 0") if keys.dtype != first_dtype: raise kvcache.CacheFormatError(f"dynamic cache layer {index} tensors must have the same dtype as layer 0") if first_dtype not in TORCH_TO_KVC: raise kvcache.CacheFormatError(f"unsupported torch cache dtype: {first_dtype}") return layers def _tensor_bytes(tensor: torch.Tensor) -> bytes: return tensor.view(torch.uint8).numpy().tobytes(order="C") def write_dynamic_cache( path: str | Path, cache: Any, *, model_revision: str, tokenizer_sha256: str, rope_theta: float, sequence_start: int = 0, ) -> dict[str, Any]: """Serialize a Transformers-style DynamicCache into a tensor-verified KVC1 file.""" layers = _cache_layers(cache) sample = layers[0][0] batch_size, kv_heads, sequence_length, head_dim = sample.shape dtype = TORCH_TO_KVC[sample.dtype] tensors: list[dict[str, Any]] = [] chunks: list[bytes] = [] offset = 0 for layer_index, pair in enumerate(layers): for kind, tensor in zip(("key", "value"), pair, strict=True): chunk = _tensor_bytes(tensor) tensors.append({ "name": f"layer.{layer_index}.{kind}", "kind": kind, "layer": layer_index, "shape": list(tensor.shape), "strides": [stride * tensor.element_size() for stride in tensor.stride()], "offset": offset, "nbytes": len(chunk), "dtype": dtype, }) chunks.append(chunk) offset += len(chunk) metadata: dict[str, Any] = { "model_revision": model_revision, "tokenizer_sha256": tokenizer_sha256, "rope_theta": rope_theta, "layers": len(layers), "kv_heads": kv_heads, "head_dim": head_dim, "dtype": dtype, "layout": "layer-major-k-then-v", "sequence_start": sequence_start, "sequence_length": sequence_length, "tensor_manifest": { "version": 1, "byte_order": "little", "batch_size": batch_size, "cache_class": "dynamic", "attention_type": "causal", "tensors": tensors, }, } kvcache.write_cache_stream(path, metadata, chunks) return metadata def _restore_tensor(record: Mapping[str, Any], payload: bytes) -> torch.Tensor: start = record["offset"] end = start + record["nbytes"] data = bytearray(payload[start:end]) return torch.frombuffer(data, dtype=KVC_TO_TORCH[record["dtype"]]).clone().reshape(record["shape"]) def read_dynamic_cache( path: str | Path, expected_identity: Mapping[str, Any] | None = None, ) -> tuple[DynamicCache, dict[str, Any]]: """Load a tensor-verified KVC1 file as a new Transformers DynamicCache.""" metadata, payload = kvcache.read_cache(path, expected_identity) manifest = metadata.get("tensor_manifest") if manifest is None: raise kvcache.CacheFormatError("a tensor_manifest is required to construct a dynamic cache") by_slot = {(record["layer"], record["kind"]): record for record in manifest["tensors"]} pairs = [] for layer in range(metadata["layers"]): pairs.append(( _restore_tensor(by_slot[(layer, "key")], payload), _restore_tensor(by_slot[(layer, "value")], payload), )) return DynamicCache(pairs), metadata