kv-cache-handoff-lab / reference /torch_kvcache.py
Brazenle's picture
feat: publish exact-model KVC1 parity gate
e77d2f1 verified
Raw History Blame Contribute Delete
5.44 kB
"""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