Nucleus-Resynthesis / runtime /src /resynthesis /none_shared_page_cache.py
Wl6adams's picture
Add portable Release 188 generation runtime
919fd68 verified
Raw
History Blame Contribute Delete
24.2 kB
"""Process-shared, authority-bound cache for immutable NoNE page weights.
This module is an external storage boundary. It may cache a reconstruction
only after the caller has authorized the accepted session/generation/manifest
and verified the immutable source object. It never stores optimizer state,
candidate state, routing state, or trainable tensor aliases.
"""
from __future__ import annotations
import fcntl
import hashlib
import json
import os
import threading
import time
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Final
import orjson
import torch
from safetensors import safe_open
from safetensors.torch import save_file
SHARED_PAGE_WEIGHTS_CACHE_SCHEMA: Final[str] = (
"nnf.resynthesis.none_shared_page_weights_cache.v1"
)
SHARED_PAGE_WEIGHTS_CACHE_LOCATOR_SCHEMA: Final[str] = (
"nnf.resynthesis.none_shared_page_weights_cache_locator.v1"
)
SHARED_PAGE_WEIGHTS_CACHE_CONTRACT_REVISION: Final[int] = 1
DEFAULT_SHARED_PAGE_WEIGHTS_CACHE_ROOT: Final[Path] = Path(
"/dev/shm/nnf-resynthesis-none-page-weights-v1"
)
_WEIGHT_NAMES: Final[tuple[str, ...]] = (
"ffn_mode_t",
"gate_t",
"up_t",
"down_t",
"glyph_down_t",
"glyph_up_t",
"translation_gate_t",
"outcome_memory_t",
"repair_memory_t",
"transfer_memory_t",
)
_ARTIFACT_NAMES: Final[frozenset[str]] = frozenset(
("page_ids_t", *_WEIGHT_NAMES)
)
_FileIdentity = tuple[int, int, int, int, int]
@dataclass(frozen=True)
class SharedNoNEPageWeights:
"""One immutable, single-page, CPU weights reconstruction."""
page_ids_t: torch.Tensor
ffn_mode_t: torch.Tensor
gate_t: torch.Tensor
up_t: torch.Tensor
down_t: torch.Tensor
glyph_down_t: torch.Tensor
glyph_up_t: torch.Tensor
translation_gate_t: torch.Tensor
outcome_memory_t: torch.Tensor
repair_memory_t: torch.Tensor
transfer_memory_t: torch.Tensor
@dataclass(frozen=True)
class SharedPageCacheAuthority:
"""Immutable source authority supplied after accepted-branch validation."""
page_id: int
source_object_sha256: str
source_object_bytes: int
source_file_identity: _FileIdentity
source_format_revision: int
materialized_dtype: torch.dtype
def external_record_boundary(self) -> dict[str, Any]:
"""Serialize cache authority only at this explicit I/O boundary."""
_validate_sha256(self.source_object_sha256)
if (
self.page_id < 0
or self.source_object_bytes < 1
or self.source_format_revision < 1
or self.source_file_identity[2] != self.source_object_bytes
):
raise RuntimeError("shared page cache source authority differs")
return {
"contractRevision": SHARED_PAGE_WEIGHTS_CACHE_CONTRACT_REVISION,
"pageId": self.page_id,
"sourceObjectSha256": self.source_object_sha256,
"sourceObjectBytes": self.source_object_bytes,
"sourceFileIdentity": list(self.source_file_identity),
"sourceFormatRevision": self.source_format_revision,
"materializedDtype": _dtype_name(self.materialized_dtype),
}
def _validate_sha256(value: object) -> None:
if not (
isinstance(value, str)
and len(value) == 64
and all(character in "0123456789abcdef" for character in value)
):
raise RuntimeError("shared page cache SHA-256 authority differs")
def _dtype_name(dtype: torch.dtype) -> str:
value = str(dtype)
if not value.startswith("torch."):
raise RuntimeError("shared page cache dtype authority differs")
return value.removeprefix("torch.")
def _canonical_json_bytes(payload: Mapping[str, Any]) -> bytes:
return json.dumps(
payload,
sort_keys=True,
separators=(",", ":"),
).encode("utf-8")
def _payload_sha256(payload: Mapping[str, Any]) -> str:
return hashlib.sha256(_canonical_json_bytes(payload)).hexdigest()
def _file_identity(path: Path) -> _FileIdentity:
status = path.stat()
return (
status.st_dev,
status.st_ino,
status.st_size,
status.st_mtime_ns,
status.st_ctime_ns,
)
def _stable_file_sha256(path: Path) -> str:
identity = _file_identity(path)
digest = hashlib.sha256()
with path.open("rb") as handle:
while block := handle.read(8 * 1024 * 1024):
digest.update(block)
if _file_identity(path) != identity:
raise RuntimeError("shared page cache file changed during hashing")
return digest.hexdigest()
def _fsync_directory(path: Path) -> None:
descriptor = os.open(path, os.O_RDONLY)
try:
os.fsync(descriptor)
finally:
os.close(descriptor)
def _protected_root(root: Path) -> Path:
root.mkdir(mode=0o700, parents=True, exist_ok=True)
root.chmod(0o700)
resolved = root.resolve()
status = resolved.stat()
if (
not resolved.is_dir()
or status.st_uid != os.geteuid()
or status.st_mode & 0o077
):
raise RuntimeError("shared page cache root is not protected")
return resolved
def _protected_regular_file(path: Path) -> None:
status = path.lstat()
if (
not path.is_file()
or path.is_symlink()
or status.st_uid != os.geteuid()
or status.st_mode & 0o077
):
raise RuntimeError("shared page cache artifact is not protected")
def _atomic_json(path: Path, payload: Mapping[str, Any]) -> None:
temporary = path.with_name(
f".{path.name}.{os.getpid()}.{threading.get_ident()}.tmp"
)
temporary.unlink(missing_ok=True)
with temporary.open("xb") as handle:
os.fchmod(handle.fileno(), 0o600)
handle.write(
orjson.dumps(
payload,
option=orjson.OPT_APPEND_NEWLINE | orjson.OPT_SORT_KEYS,
)
)
handle.flush()
os.fsync(handle.fileno())
os.replace(temporary, path)
path.chmod(0o600)
_fsync_directory(path.parent)
def _record_with_payload_sha256(payload: Mapping[str, Any]) -> dict[str, Any]:
record = dict(payload)
record["recordPayloadSha256"] = _payload_sha256(record)
return record
def _validated_record(path: Path, *, schema: str) -> dict[str, Any]:
_protected_regular_file(path)
try:
raw = json.loads(path.read_text(encoding="utf-8"))
except (OSError, UnicodeDecodeError, json.JSONDecodeError) as error:
raise RuntimeError("shared page cache record is malformed") from error
if not isinstance(raw, dict):
raise RuntimeError("shared page cache record is malformed")
recorded_sha256 = raw.get("recordPayloadSha256")
unsigned = dict(raw)
unsigned.pop("recordPayloadSha256", None)
if (
raw.get("schema") != schema
or not isinstance(recorded_sha256, str)
or _payload_sha256(unsigned) != recorded_sha256
):
raise RuntimeError("shared page cache record authority differs")
return raw
def _weights_tensors(
weights: SharedNoNEPageWeights,
) -> tuple[tuple[str, torch.Tensor], ...]:
return (
("page_ids_t", weights.page_ids_t),
*((name, getattr(weights, name)) for name in _WEIGHT_NAMES),
)
def _validate_weights(
weights: SharedNoNEPageWeights,
*,
authority: SharedPageCacheAuthority,
) -> None:
tensors = _weights_tensors(weights)
if (
weights.page_ids_t.shape != (1,)
or weights.page_ids_t.dtype != torch.long
or int(weights.page_ids_t.reshape(())) != authority.page_id
or any(tensor.device.type != "cpu" for _name, tensor in tensors)
or any(tensor.requires_grad for _name, tensor in tensors)
or any(tensor.shape[0] != 1 for _name, tensor in tensors)
or any(
tensor.dtype != authority.materialized_dtype
for name, tensor in tensors
if name != "page_ids_t"
)
):
raise RuntimeError("shared page cache tensor authority differs")
if any(
not torch.isfinite(tensor).all()
for name, tensor in tensors
if name != "page_ids_t"
):
raise RuntimeError("shared page cache contains nonfinite weights")
def _geometry_record(
weights: SharedNoNEPageWeights,
) -> dict[str, dict[str, Any]]:
return {
name: {
"shape": list(tensor.shape),
"dtype": _dtype_name(tensor.dtype),
}
for name, tensor in _weights_tensors(weights)
}
def _private_weights(weights: SharedNoNEPageWeights) -> SharedNoNEPageWeights:
values = {
name: tensor.detach().to(device="cpu", copy=True).contiguous()
for name, tensor in _weights_tensors(weights)
}
return SharedNoNEPageWeights(**values)
def _artifact_payload(
weights: SharedNoNEPageWeights,
) -> dict[str, torch.Tensor]:
return {
name: tensor.detach().to(device="cpu", copy=True).contiguous()
for name, tensor in _weights_tensors(weights)
}
class SharedPageWeightsCache:
"""Bounded process-shared cache with per-authority single-flight."""
def __init__(
self,
root: Path = DEFAULT_SHARED_PAGE_WEIGHTS_CACHE_ROOT,
*,
budget_bytes: int | None = None,
) -> None:
self.root = _protected_root(root)
physical_memory_bytes = (
os.sysconf("SC_PHYS_PAGES") * os.sysconf("SC_PAGE_SIZE")
)
filesystem_bytes = (
os.statvfs(self.root).f_blocks * os.statvfs(self.root).f_frsize
)
computed_budget = max(
os.sysconf("SC_PAGE_SIZE"),
min(physical_memory_bytes // 16, filesystem_bytes // 4),
)
self.budget_bytes = (
computed_budget if budget_bytes is None else budget_bytes
)
if self.budget_bytes < 1:
raise ValueError("shared page cache budget must be positive")
@staticmethod
def authority_prefix(authority: SharedPageCacheAuthority) -> str:
return _payload_sha256(authority.external_record_boundary())
def _lock_path(self, authority_prefix: str) -> Path:
_validate_sha256(authority_prefix)
return self.root / f".{authority_prefix}.lock"
def _locator_path(self, authority_prefix: str) -> Path:
return self.root / f"{authority_prefix}.locator.json"
def _artifact_path(self, authority_prefix: str, cache_key: str) -> Path:
_validate_sha256(cache_key)
return self.root / f"{authority_prefix}.{cache_key}.safetensors"
def _receipt_path(self, authority_prefix: str, cache_key: str) -> Path:
_validate_sha256(cache_key)
return self.root / f"{authority_prefix}.{cache_key}.receipt.json"
def _open_lock(self, authority_prefix: str) -> Any:
lock_path = self._lock_path(authority_prefix)
descriptor = os.open(
lock_path,
os.O_CREAT | os.O_RDWR | os.O_NOFOLLOW,
0o600,
)
os.fchmod(descriptor, 0o600)
return os.fdopen(descriptor, "a+b")
def _load_locator(
self,
authority: SharedPageCacheAuthority,
authority_prefix: str,
) -> dict[str, Any] | None:
locator_path = self._locator_path(authority_prefix)
if not locator_path.exists():
return None
locator = _validated_record(
locator_path,
schema=SHARED_PAGE_WEIGHTS_CACHE_LOCATOR_SCHEMA,
)
cache_key = locator.get("cacheKey")
receipt_sha256 = locator.get("receiptSha256")
if (
locator.get("authority") != authority.external_record_boundary()
or locator.get("authorityPrefix") != authority_prefix
):
raise RuntimeError("shared page cache locator authority differs")
_validate_sha256(cache_key)
_validate_sha256(receipt_sha256)
return locator
def _load_hit(
self,
*,
authority: SharedPageCacheAuthority,
authority_prefix: str,
locator: Mapping[str, Any],
) -> SharedNoNEPageWeights | None:
cache_key = str(locator["cacheKey"])
artifact_path = self._artifact_path(authority_prefix, cache_key)
receipt_path = self._receipt_path(authority_prefix, cache_key)
if not artifact_path.exists() or not receipt_path.exists():
return None
if _stable_file_sha256(receipt_path) != locator["receiptSha256"]:
raise RuntimeError("shared page cache receipt hash differs")
receipt = _validated_record(
receipt_path,
schema=SHARED_PAGE_WEIGHTS_CACHE_SCHEMA,
)
artifact = receipt.get("artifact")
geometry = receipt.get("geometry")
if (
receipt.get("authority") != authority.external_record_boundary()
or receipt.get("authorityPrefix") != authority_prefix
or receipt.get("cacheKey") != cache_key
or receipt.get("weightsOnly") is not True
or receipt.get("optimizerStateStored") is not False
or receipt.get("candidateStateStored") is not False
or not isinstance(artifact, dict)
or artifact.get("path") != artifact_path.name
or not isinstance(artifact.get("bytes"), int)
or not isinstance(geometry, dict)
):
raise RuntimeError("shared page cache receipt authority differs")
artifact_sha256 = artifact.get("sha256")
geometry_sha256 = receipt.get("geometrySha256")
_validate_sha256(artifact_sha256)
_validate_sha256(geometry_sha256)
_protected_regular_file(artifact_path)
if (
artifact_path.stat().st_size != artifact["bytes"]
or _stable_file_sha256(artifact_path) != artifact_sha256
):
raise RuntimeError("shared page cache artifact hash differs")
identity = _file_identity(artifact_path)
with safe_open( # type: ignore[no-untyped-call]
str(artifact_path),
framework="pt",
device="cpu",
) as handle:
if set(handle.keys()) != _ARTIFACT_NAMES:
raise RuntimeError("shared page cache tensor set differs")
loaded = SharedNoNEPageWeights(
**{
name: handle.get_tensor(name)
for name in ("page_ids_t", *_WEIGHT_NAMES)
}
)
if _file_identity(artifact_path) != identity:
raise RuntimeError(
"shared page cache artifact changed during materialization"
)
_validate_weights(loaded, authority=authority)
rebuilt_geometry = _geometry_record(loaded)
if (
rebuilt_geometry != geometry
or _payload_sha256(rebuilt_geometry) != geometry_sha256
):
raise RuntimeError("shared page cache geometry differs")
key_payload = {
"authority": authority.external_record_boundary(),
"artifactSha256": artifact_sha256,
"geometrySha256": geometry_sha256,
}
if _payload_sha256(key_payload) != cache_key:
raise RuntimeError("shared page cache content key differs")
os.utime(self._locator_path(authority_prefix), None)
return _private_weights(loaded)
def _publish(
self,
*,
authority: SharedPageCacheAuthority,
authority_prefix: str,
weights: SharedNoNEPageWeights,
) -> SharedNoNEPageWeights:
_validate_weights(weights, authority=authority)
geometry = _geometry_record(weights)
geometry_sha256 = _payload_sha256(geometry)
temporary = self.root / (
f".{authority_prefix}.{os.getpid()}."
f"{threading.get_ident()}.{time.time_ns()}.tmp"
)
save_file(_artifact_payload(weights), str(temporary))
temporary.chmod(0o600)
with temporary.open("rb") as handle:
os.fsync(handle.fileno())
artifact_sha256 = _stable_file_sha256(temporary)
key_payload = {
"authority": authority.external_record_boundary(),
"artifactSha256": artifact_sha256,
"geometrySha256": geometry_sha256,
}
cache_key = _payload_sha256(key_payload)
artifact_path = self._artifact_path(authority_prefix, cache_key)
receipt_path = self._receipt_path(authority_prefix, cache_key)
locator_path = self._locator_path(authority_prefix)
pair_exists = artifact_path.exists() and receipt_path.exists()
if (artifact_path.exists() or receipt_path.exists()) and not pair_exists:
artifact_path.unlink(missing_ok=True)
receipt_path.unlink(missing_ok=True)
_fsync_directory(self.root)
if pair_exists:
temporary.unlink(missing_ok=True)
else:
os.replace(temporary, artifact_path)
artifact_path.chmod(0o400)
_fsync_directory(self.root)
receipt = _record_with_payload_sha256(
{
"schema": SHARED_PAGE_WEIGHTS_CACHE_SCHEMA,
"authority": authority.external_record_boundary(),
"authorityPrefix": authority_prefix,
"cacheKey": cache_key,
"geometry": geometry,
"geometrySha256": geometry_sha256,
"artifact": {
"path": artifact_path.name,
"sha256": artifact_sha256,
"bytes": artifact_path.stat().st_size,
},
"weightsOnly": True,
"optimizerStateStored": False,
"candidateStateStored": False,
"acceptedSourceAuthorityRequiredOnEveryAccess": True,
}
)
_atomic_json(receipt_path, receipt)
receipt_sha256 = _stable_file_sha256(receipt_path)
locator = _record_with_payload_sha256(
{
"schema": SHARED_PAGE_WEIGHTS_CACHE_LOCATOR_SCHEMA,
"authority": authority.external_record_boundary(),
"authorityPrefix": authority_prefix,
"cacheKey": cache_key,
"receiptSha256": receipt_sha256,
}
)
_atomic_json(locator_path, locator)
loaded = self._load_hit(
authority=authority,
authority_prefix=authority_prefix,
locator=locator,
)
if loaded is None:
raise RuntimeError("shared page cache publication is incomplete")
return loaded
def load_or_build(
self,
*,
authority: SharedPageCacheAuthority,
builder: Callable[[], SharedNoNEPageWeights],
) -> SharedNoNEPageWeights:
"""Load or publish one exact weight row under a per-key process lock."""
authority_prefix = self.authority_prefix(authority)
with self._open_lock(authority_prefix) as handle:
fcntl.flock(handle.fileno(), fcntl.LOCK_EX)
try:
locator = self._load_locator(authority, authority_prefix)
loaded = (
self._load_hit(
authority=authority,
authority_prefix=authority_prefix,
locator=locator,
)
if locator is not None
else None
)
if loaded is None:
if locator is not None:
stale_cache_key = locator.get("cacheKey")
_validate_sha256(stale_cache_key)
assert isinstance(stale_cache_key, str)
self._locator_path(authority_prefix).unlink(
missing_ok=True
)
self._receipt_path(
authority_prefix,
stale_cache_key,
).unlink(missing_ok=True)
self._artifact_path(
authority_prefix,
stale_cache_key,
).unlink(missing_ok=True)
_fsync_directory(self.root)
loaded = self._publish(
authority=authority,
authority_prefix=authority_prefix,
weights=builder(),
)
finally:
fcntl.flock(handle.fileno(), fcntl.LOCK_UN)
# Eviction owns a separate global lock and begins only after this key's
# single-flight lock is released. Returned tensors are private copies.
self._evict_to_budget()
return loaded
def _evict_to_budget(self) -> None:
eviction_path = self.root / ".eviction.lock"
descriptor = os.open(
eviction_path,
os.O_CREAT | os.O_RDWR | os.O_NOFOLLOW,
0o600,
)
os.fchmod(descriptor, 0o600)
with os.fdopen(descriptor, "a+b") as eviction_handle:
fcntl.flock(eviction_handle.fileno(), fcntl.LOCK_EX)
try:
artifacts = tuple(self.root.glob("*.safetensors"))
resident_bytes = sum(path.stat().st_size for path in artifacts)
if resident_bytes <= self.budget_bytes:
return
locators = sorted(
self.root.glob("*.locator.json"),
key=lambda path: path.stat().st_mtime_ns,
)
for locator_path in locators:
if resident_bytes <= self.budget_bytes:
break
authority_prefix = locator_path.name.removesuffix(
".locator.json"
)
try:
lock_handle = self._open_lock(authority_prefix)
except (OSError, RuntimeError):
continue
with lock_handle:
try:
fcntl.flock(
lock_handle.fileno(),
fcntl.LOCK_EX | fcntl.LOCK_NB,
)
except BlockingIOError:
continue
try:
locator = _validated_record(
locator_path,
schema=(
SHARED_PAGE_WEIGHTS_CACHE_LOCATOR_SCHEMA
),
)
cache_key = locator.get("cacheKey")
_validate_sha256(cache_key)
assert isinstance(cache_key, str)
artifact_path = self._artifact_path(
authority_prefix,
cache_key,
)
receipt_path = self._receipt_path(
authority_prefix,
cache_key,
)
artifact_bytes = (
artifact_path.stat().st_size
if artifact_path.is_file()
else 0
)
locator_path.unlink(missing_ok=True)
receipt_path.unlink(missing_ok=True)
artifact_path.unlink(missing_ok=True)
resident_bytes -= artifact_bytes
_fsync_directory(self.root)
finally:
fcntl.flock(
lock_handle.fileno(),
fcntl.LOCK_UN,
)
finally:
fcntl.flock(eviction_handle.fileno(), fcntl.LOCK_UN)