Download reference/kvcache.py from Brazenle/kv-cache-handoff-lab: direct link, hf CLI and curl.
- Browser
- Download file 15.1 kB
-
https://huggingface.co/Brazenle/kv-cache-handoff-lab/resolve/main/reference/kvcache.py
- Command line
-
hf download hf://Brazenle/kv-cache-handoff-lab/reference/kvcache.py
-
curl -L -o kvcache.py https://huggingface.co/Brazenle/kv-cache-handoff-lab/resolve/main/reference/kvcache.py
15.1 kB
| """Portable, fail-closed raw KV-cache container. | |
| KVC1 stores already-materialized key/value tensor bytes plus the exact model, | |
| tokenizer, tensor geometry, RoPE, dtype, layout, and sequence-position identity | |
| needed to decide whether a runtime may safely reuse them. Cross-architecture | |
| translation is deliberately out of scope for this byte container. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import hashlib | |
| import hmac | |
| import json | |
| import os | |
| from pathlib import Path | |
| import struct | |
| import tempfile | |
| from collections.abc import Iterable | |
| from typing import Any, Mapping | |
| MAGIC = b"KVC1" | |
| PREFIX = struct.Struct(">4sIQ32s") | |
| REQUIRED_FIELDS = { | |
| "model_revision", | |
| "tokenizer_sha256", | |
| "rope_theta", | |
| "layers", | |
| "kv_heads", | |
| "head_dim", | |
| "dtype", | |
| "layout", | |
| "sequence_start", | |
| "sequence_length", | |
| } | |
| OPTIONAL_FIELDS = {"tensor_manifest"} | |
| KNOWN_FIELDS = REQUIRED_FIELDS | OPTIONAL_FIELDS | |
| ALLOWED_DTYPES = {"f16", "bf16", "f32", "i8", "u8"} | |
| DTYPE_BYTES = {"f16": 2, "bf16": 2, "f32": 4, "i8": 1, "u8": 1} | |
| ALLOWED_LAYOUTS = {"layer-major-k-then-v"} | |
| MAX_METADATA_BYTES = 1_048_576 | |
| MANIFEST_FIELDS = {"version", "byte_order", "batch_size", "cache_class", "attention_type", "tensors"} | |
| TENSOR_FIELDS = {"name", "kind", "layer", "shape", "strides", "offset", "nbytes", "dtype"} | |
| class CacheFormatError(ValueError): | |
| """The container is malformed, incomplete, corrupt, or unsupported.""" | |
| class CacheCompatibilityError(ValueError): | |
| """The container is valid but does not match the requested runtime.""" | |
| def _exact_fields(value: Mapping[str, Any], expected: set[str], label: str) -> None: | |
| missing = expected.difference(value) | |
| extra = set(value).difference(expected) | |
| if missing or extra: | |
| raise CacheFormatError(f"{label} fields differ: missing={sorted(missing)}, extra={sorted(extra)}") | |
| def _positive_integer(value: Any, label: str) -> int: | |
| if not isinstance(value, int) or isinstance(value, bool) or value <= 0: | |
| raise CacheFormatError(f"{label} must be a positive integer") | |
| return value | |
| def _contiguous_byte_strides(shape: list[int], item_bytes: int) -> list[int]: | |
| strides = [item_bytes] * len(shape) | |
| for index in range(len(shape) - 2, -1, -1): | |
| strides[index] = strides[index + 1] * shape[index + 1] | |
| return strides | |
| def _validate_tensor_manifest(metadata: Mapping[str, Any], manifest: Any) -> int: | |
| if not isinstance(manifest, Mapping): | |
| raise CacheFormatError("tensor_manifest must be a mapping") | |
| manifest = dict(manifest) | |
| _exact_fields(manifest, MANIFEST_FIELDS, "tensor_manifest") | |
| if manifest["version"] != 1: | |
| raise CacheFormatError("tensor_manifest version must be 1") | |
| if manifest["byte_order"] not in {"little", "big"}: | |
| raise CacheFormatError("tensor_manifest byte_order must be little or big") | |
| batch_size = _positive_integer(manifest["batch_size"], "tensor_manifest batch_size") | |
| for field in ("cache_class", "attention_type"): | |
| if not isinstance(manifest[field], str) or not manifest[field].strip(): | |
| raise CacheFormatError(f"tensor_manifest {field} must be a non-empty string") | |
| tensors = manifest["tensors"] | |
| if not isinstance(tensors, list): | |
| raise CacheFormatError("tensor_manifest tensors must be a list") | |
| expected_count = metadata["layers"] * 2 | |
| if len(tensors) != expected_count: | |
| raise CacheFormatError(f"tensor_manifest must contain exactly {expected_count} key/value tensors") | |
| expected_shape = [batch_size, metadata["kv_heads"], metadata["sequence_length"], metadata["head_dim"]] | |
| expected_offset = 0 | |
| seen_names: set[str] = set() | |
| seen_slots: set[tuple[int, str]] = set() | |
| for index, raw_tensor in enumerate(tensors): | |
| if not isinstance(raw_tensor, Mapping): | |
| raise CacheFormatError(f"tensor_manifest tensor {index} must be a mapping") | |
| tensor = dict(raw_tensor) | |
| _exact_fields(tensor, TENSOR_FIELDS, f"tensor_manifest tensor {index}") | |
| name = tensor["name"] | |
| if not isinstance(name, str) or not name: | |
| raise CacheFormatError(f"tensor_manifest tensor {index} name must be non-empty") | |
| if name in seen_names: | |
| raise CacheFormatError(f"tensor_manifest tensor name is duplicated: {name}") | |
| seen_names.add(name) | |
| kind = tensor["kind"] | |
| if kind not in {"key", "value"}: | |
| raise CacheFormatError(f"tensor_manifest tensor {index} kind must be key or value") | |
| layer = tensor["layer"] | |
| if not isinstance(layer, int) or isinstance(layer, bool) or not 0 <= layer < metadata["layers"]: | |
| raise CacheFormatError(f"tensor_manifest tensor {index} layer is out of range") | |
| slot = (layer, kind) | |
| if slot in seen_slots: | |
| raise CacheFormatError(f"tensor_manifest duplicates layer {layer} {kind}") | |
| seen_slots.add(slot) | |
| shape = tensor["shape"] | |
| if shape != expected_shape: | |
| raise CacheFormatError(f"tensor_manifest tensor {index} shape must equal {expected_shape}") | |
| dtype = tensor["dtype"] | |
| if dtype != metadata["dtype"]: | |
| raise CacheFormatError(f"tensor_manifest tensor {index} dtype must match metadata dtype") | |
| expected_nbytes = DTYPE_BYTES[dtype] | |
| for dimension in shape: | |
| expected_nbytes *= _positive_integer(dimension, f"tensor_manifest tensor {index} shape dimension") | |
| if tensor["nbytes"] != expected_nbytes: | |
| raise CacheFormatError(f"tensor_manifest tensor {index} nbytes must equal {expected_nbytes}") | |
| expected_strides = _contiguous_byte_strides(shape, DTYPE_BYTES[dtype]) | |
| if tensor["strides"] != expected_strides: | |
| raise CacheFormatError( | |
| f"tensor_manifest tensor {index} must use C-contiguous byte strides {expected_strides}" | |
| ) | |
| if tensor["offset"] != expected_offset: | |
| raise CacheFormatError( | |
| f"tensor_manifest tensor {index} offsets must be contiguous; expected {expected_offset}" | |
| ) | |
| expected_offset += expected_nbytes | |
| expected_slots = {(layer, kind) for layer in range(metadata["layers"]) for kind in ("key", "value")} | |
| if seen_slots != expected_slots: | |
| raise CacheFormatError("tensor_manifest must contain one key and one value tensor for every layer") | |
| return expected_offset | |
| def _validate_payload_length(metadata: Mapping[str, Any], payload_length: int) -> None: | |
| manifest = metadata.get("tensor_manifest") | |
| if manifest is None: | |
| return | |
| expected = _validate_tensor_manifest(metadata, manifest) | |
| if payload_length != expected: | |
| raise CacheFormatError( | |
| f"tensor manifest declares {expected} payload bytes but container has {payload_length}" | |
| ) | |
| def _plain_dict(metadata: Mapping[str, Any]) -> dict[str, Any]: | |
| if not isinstance(metadata, Mapping): | |
| raise CacheFormatError("metadata must be a mapping") | |
| value = dict(metadata) | |
| missing = REQUIRED_FIELDS.difference(value) | |
| extra = set(value).difference(KNOWN_FIELDS) | |
| if missing or extra: | |
| raise CacheFormatError(f"metadata fields differ: missing={sorted(missing)}, extra={sorted(extra)}") | |
| if not isinstance(value["model_revision"], str) or "@" not in value["model_revision"]: | |
| raise CacheFormatError("model_revision must identify an immutable revision") | |
| tokenizer_hash = value["tokenizer_sha256"] | |
| if not isinstance(tokenizer_hash, str) or len(tokenizer_hash) != 64: | |
| raise CacheFormatError("tokenizer_sha256 must contain 64 hexadecimal characters") | |
| try: | |
| int(tokenizer_hash, 16) | |
| except ValueError as error: | |
| raise CacheFormatError("tokenizer_sha256 is not hexadecimal") from error | |
| if not isinstance(value["rope_theta"], (int, float)) or isinstance(value["rope_theta"], bool) or value["rope_theta"] <= 0: | |
| raise CacheFormatError("rope_theta must be positive") | |
| for field in ("layers", "kv_heads", "head_dim", "sequence_length"): | |
| if not isinstance(value[field], int) or isinstance(value[field], bool) or value[field] <= 0: | |
| raise CacheFormatError(f"{field} must be a positive integer") | |
| if not isinstance(value["sequence_start"], int) or isinstance(value["sequence_start"], bool) or value["sequence_start"] < 0: | |
| raise CacheFormatError("sequence_start must be a non-negative integer") | |
| if value["dtype"] not in ALLOWED_DTYPES: | |
| raise CacheFormatError("unsupported dtype") | |
| if value["layout"] not in ALLOWED_LAYOUTS: | |
| raise CacheFormatError("unsupported layout") | |
| if "tensor_manifest" in value: | |
| _validate_tensor_manifest(value, value["tensor_manifest"]) | |
| return value | |
| def _metadata_bytes(metadata: Mapping[str, Any]) -> bytes: | |
| try: | |
| encoded = json.dumps(_plain_dict(metadata), sort_keys=True, separators=(",", ":"), allow_nan=False).encode("utf-8") | |
| except (TypeError, ValueError) as error: | |
| if isinstance(error, CacheFormatError): | |
| raise | |
| raise CacheFormatError("metadata is not canonical JSON") from error | |
| if len(encoded) > MAX_METADATA_BYTES: | |
| raise CacheFormatError("metadata is too large") | |
| return encoded | |
| def write_cache(path: str | os.PathLike[str], metadata: Mapping[str, Any], payload: bytes) -> None: | |
| """Atomically publish one KVC1 generation.""" | |
| if type(payload) is not bytes: | |
| raise CacheFormatError("payload must be raw bytes") | |
| return write_cache_stream(Path(path), metadata, (payload,)) | |
| def write_cache_stream(path: str | os.PathLike[str], metadata: Mapping[str, Any], chunks: Iterable[bytes | bytearray | memoryview]) -> None: | |
| destination = Path(path) | |
| metadata_dict = _plain_dict(metadata) | |
| metadata_bytes = _metadata_bytes(metadata_dict) | |
| destination.parent.mkdir(parents=True, exist_ok=True) | |
| temporary_name: str | None = None | |
| try: | |
| with tempfile.NamedTemporaryFile( | |
| mode="wb", | |
| prefix=f".{destination.name}.", | |
| suffix=".tmp", | |
| dir=destination.parent, | |
| delete=False, | |
| ) as temporary: | |
| temporary_name = temporary.name | |
| placeholder = PREFIX.pack(MAGIC, len(metadata_bytes), 0, b"\x00" * 32) | |
| temporary.write(placeholder) | |
| temporary.write(metadata_bytes) | |
| sha = hashlib.sha256() | |
| length = 0 | |
| for chunk in chunks: | |
| if not isinstance(chunk, (bytes, bytearray, memoryview)): | |
| raise CacheFormatError("chunk must be bytes, bytearray, or memoryview") | |
| chunk_bytes = bytes(chunk) | |
| temporary.write(chunk_bytes) | |
| sha.update(chunk_bytes) | |
| length += len(chunk_bytes) | |
| _validate_payload_length(metadata_dict, length) | |
| digest = sha.digest() | |
| temporary.seek(0) | |
| temporary.write(PREFIX.pack(MAGIC, len(metadata_bytes), length, digest)) | |
| temporary.flush() | |
| os.fsync(temporary.fileno()) | |
| os.replace(temporary_name, destination) | |
| temporary_name = None | |
| finally: | |
| if temporary_name is not None: | |
| try: | |
| os.unlink(temporary_name) | |
| except FileNotFoundError: | |
| pass | |
| def _decode(path: str | os.PathLike[str]) -> tuple[dict[str, Any], bytes, str]: | |
| try: | |
| raw = Path(path).read_bytes() | |
| except OSError as error: | |
| raise CacheFormatError(f"cache could not be read: {error}") from error | |
| if len(raw) < PREFIX.size: | |
| raise CacheFormatError("container is truncated") | |
| try: | |
| magic, metadata_length, payload_length, expected_digest = PREFIX.unpack_from(raw) | |
| except struct.error as error: | |
| raise CacheFormatError("container prefix is malformed") from error | |
| if magic != MAGIC: | |
| raise CacheFormatError("unsupported container magic or version") | |
| if metadata_length == 0 or metadata_length > MAX_METADATA_BYTES: | |
| raise CacheFormatError("metadata length is invalid") | |
| expected_length = PREFIX.size + metadata_length + payload_length | |
| if len(raw) != expected_length: | |
| raise CacheFormatError("container length does not match its header") | |
| metadata_raw = raw[PREFIX.size:PREFIX.size + metadata_length] | |
| payload = raw[PREFIX.size + metadata_length:] | |
| try: | |
| decoded = json.loads(metadata_raw.decode("utf-8")) | |
| except (UnicodeDecodeError, json.JSONDecodeError) as error: | |
| raise CacheFormatError("metadata is not valid UTF-8 JSON") from error | |
| metadata = _plain_dict(decoded) | |
| if _metadata_bytes(metadata) != metadata_raw: | |
| raise CacheFormatError("metadata is not in canonical form") | |
| actual_digest = hashlib.sha256(payload).digest() | |
| if not hmac.compare_digest(actual_digest, expected_digest): | |
| raise CacheFormatError("payload checksum mismatch") | |
| _validate_payload_length(metadata, len(payload)) | |
| return metadata, payload, actual_digest.hex() | |
| def read_cache(path: str | os.PathLike[str], expected_identity: Mapping[str, Any] | None = None) -> tuple[dict[str, Any], bytes]: | |
| metadata, payload, _ = _decode(path) | |
| if expected_identity is not None: | |
| if not isinstance(expected_identity, Mapping): | |
| raise CacheCompatibilityError("expected_identity must be a mapping") | |
| expected = dict(expected_identity) | |
| unknown = set(expected).difference(KNOWN_FIELDS) | |
| if unknown: | |
| raise CacheCompatibilityError(f"expected_identity has unknown fields: {sorted(unknown)}") | |
| try: | |
| _plain_dict({**metadata, **expected}) | |
| except CacheFormatError as error: | |
| raise CacheCompatibilityError(f"expected_identity is invalid: {error}") from error | |
| differences = [field for field in sorted(expected) if metadata[field] != expected[field]] | |
| if differences: | |
| raise CacheCompatibilityError(f"cache is incompatible: {', '.join(differences)}") | |
| return metadata, payload | |
| def inspect_cache(path: str | os.PathLike[str]) -> dict[str, Any]: | |
| metadata, payload, digest = _decode(path) | |
| manifest = metadata.get("tensor_manifest") | |
| receipt = { | |
| "format": "KVC1", | |
| "metadata": metadata, | |
| "payload_bytes": len(payload), | |
| "payload_sha256": digest, | |
| } | |
| if manifest is not None: | |
| receipt["tensor_manifest_verified"] = True | |
| receipt["tensor_count"] = len(manifest["tensors"]) | |
| return receipt | |
| def main() -> int: | |
| parser = argparse.ArgumentParser(description="Inspect a portable raw KV-cache container.") | |
| subparsers = parser.add_subparsers(dest="command", required=True) | |
| inspect_parser = subparsers.add_parser("inspect") | |
| inspect_parser.add_argument("path") | |
| args = parser.parse_args() | |
| if args.command == "inspect": | |
| try: | |
| print(json.dumps(inspect_cache(args.path), sort_keys=True)) | |
| except (CacheFormatError, CacheCompatibilityError) as error: | |
| parser.exit(1, f"kvcache: {error}\n") | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |