Brazenle's picture
Add accepted KVC1 reference implementation
2c76aad verified
Raw
History Blame Contribute Delete
7.78 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 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",
}
ALLOWED_DTYPES = {"f16", "bf16", "f32", "i8", "u8"}
ALLOWED_LAYOUTS = {"layer-major-k-then-v"}
MAX_METADATA_BYTES = 1_048_576
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 _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(REQUIRED_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")
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")
destination = Path(path)
destination.parent.mkdir(parents=True, exist_ok=True)
metadata_bytes = _metadata_bytes(metadata)
digest = hashlib.sha256(payload).digest()
prefix = PREFIX.pack(MAGIC, len(metadata_bytes), len(payload), digest)
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
temporary.write(prefix)
temporary.write(metadata_bytes)
temporary.write(payload)
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")
return metadata, payload, actual_digest.hex()
def read_cache(path: str | os.PathLike[str], expected_metadata: Mapping[str, Any] | None = None) -> tuple[dict[str, Any], bytes]:
metadata, payload, _ = _decode(path)
if expected_metadata is not None:
expected = _plain_dict(expected_metadata)
differences = [field for field in sorted(REQUIRED_FIELDS) 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)
return {
"format": "KVC1",
"metadata": metadata,
"payload_bytes": len(payload),
"payload_sha256": digest,
}
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())