Download reference/torch_kvcache.py from Brazenle/kv-cache-handoff-lab: direct link, hf CLI and curl.
- Browser
- Download file 5.44 kB
-
https://huggingface.co/Brazenle/kv-cache-handoff-lab/resolve/main/reference/torch_kvcache.py
- Command line
-
hf download hf://Brazenle/kv-cache-handoff-lab/reference/torch_kvcache.py
-
curl -L -o torch_kvcache.py https://huggingface.co/Brazenle/kv-cache-handoff-lab/resolve/main/reference/torch_kvcache.py
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 | |