cortex.6.sol / model /init.py
openhands
openhands
feat(cortex): add CORTEX training pipeline and model audit
6fbe100
Raw History Blame Contribute Delete
7.54 kB
"""CORTEX initialisation and checkpoint I/O.
Creates brand-new CORTEX weights from scratch and writes them as self-describing
checkpoints that record exactly how they were produced.
It never reads, writes or modifies the distributed 1.65T checkpoint that this
repository also hosts. New CORTEX checkpoints are separate artefacts with their own
provenance record.
"""
from __future__ import annotations
import hashlib
import json
import platform
import subprocess
import time
from pathlib import Path
import torch
from model.cortex_model import CortexConfig, CortexForCausalLM, count_parameters
__all__ = [
"init_cortex_model",
"save_cortex_checkpoint",
"load_cortex_checkpoint",
"sha256_file",
"checkpoint_provenance",
]
CHECKPOINT_FORMAT = "cortex-checkpoint-v1"
def _git_commit() -> str | None:
try:
out = subprocess.run(
["git", "rev-parse", "HEAD"],
cwd=Path(__file__).resolve().parent,
capture_output=True, text=True, timeout=10,
)
return out.stdout.strip() or None
except Exception:
return None
def checkpoint_provenance(config: CortexConfig, step: int, extra: dict | None = None) -> dict:
"""Describe how a checkpoint was produced. Recorded inside every checkpoint."""
record = {
"format": CHECKPOINT_FORMAT,
"model_name": config.model_name,
"developer": "Frankenstein-Labs",
"weights": "initialised and trained by Frankenstein-Labs",
"base_model": None,
"base_model_note": (
"No parent model. These weights are not derived from any other model, so no "
"base_model attribution applies."
),
"step": step,
"created_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
"torch_version": torch.__version__,
"python_version": platform.python_version(),
"source_commit": _git_commit(),
"license": "mit",
}
if extra:
record.update(extra)
return record
def init_cortex_model(config: CortexConfig, seed: int | None = None) -> CortexForCausalLM:
"""Create a new CORTEX model with randomly initialised weights."""
if seed is not None:
torch.manual_seed(seed)
model = CortexForCausalLM(config)
model.eval()
return model
def sha256_file(path: str | Path, chunk: int = 1 << 20) -> str:
h = hashlib.sha256()
with open(path, "rb") as fh:
while block := fh.read(chunk):
h.update(block)
return h.hexdigest()
def _split_tied_tensors(state: dict) -> tuple[dict, dict]:
"""Separate storage-shared tensors so safetensors can write the checkpoint.
``safetensors`` refuses to serialise two keys that point at the same storage, and a
tied embedding does exactly that. Only one copy is written; the aliases are recorded
and re-expanded on load, so the checkpoint stays lossless without storing the tied
matrix twice.
"""
unique: dict = {}
aliases: dict[str, str] = {}
by_ptr: dict[int, str] = {}
for name, tensor in state.items():
ptr = tensor.data_ptr()
if ptr in by_ptr:
aliases[name] = by_ptr[ptr]
else:
by_ptr[ptr] = name
unique[name] = tensor
return unique, aliases
def save_cortex_checkpoint(
model: CortexForCausalLM,
config: CortexConfig,
out_dir: str | Path,
step: int = 0,
optimizer: torch.optim.Optimizer | None = None,
extra_provenance: dict | None = None,
) -> Path:
"""Write a CORTEX checkpoint directory.
Layout::
<out_dir>/
config.json model hyper-parameters
weights.safetensors
provenance.json how the weights were produced
optimizer.pt only when an optimizer is passed
manifest.json sizes and SHA-256 of the files above
"""
out_dir = Path(out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
config.to_json(out_dir / "config.json")
state = {k: v.detach().cpu() for k, v in model.state_dict().items()}
unique, aliases = _split_tied_tensors(state)
weights_path = out_dir / "weights.safetensors"
try:
from safetensors.torch import save_file
save_file(unique, str(weights_path))
except ImportError: # pragma: no cover - fallback when safetensors is absent
weights_path = out_dir / "weights.pt"
torch.save(state, weights_path)
provenance = checkpoint_provenance(config, step, extra_provenance)
provenance["parameter_count"] = count_parameters(model)
provenance["tensor_count"] = len(state)
provenance["stored_tensors"] = len(unique)
provenance["tied_tensors"] = aliases
(out_dir / "provenance.json").write_text(
json.dumps(provenance, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
)
if optimizer is not None:
torch.save(optimizer.state_dict(), out_dir / "optimizer.pt")
manifest = {"format": CHECKPOINT_FORMAT, "files": {}}
for name in sorted(p.name for p in out_dir.iterdir() if p.name != "manifest.json"):
f = out_dir / name
manifest["files"][name] = {"bytes": f.stat().st_size, "sha256": sha256_file(f)}
(out_dir / "manifest.json").write_text(
json.dumps(manifest, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
)
return out_dir
def load_cortex_checkpoint(
ckpt_dir: str | Path, strict: bool = True
) -> tuple[CortexForCausalLM, CortexConfig, dict]:
"""Load a CORTEX checkpoint, verifying its manifest first."""
ckpt_dir = Path(ckpt_dir)
manifest_path = ckpt_dir / "manifest.json"
if not manifest_path.exists():
raise FileNotFoundError(f"no manifest.json in {ckpt_dir}; not a CORTEX checkpoint")
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
if manifest.get("format") != CHECKPOINT_FORMAT:
raise ValueError(f"unknown checkpoint format: {manifest.get('format')!r}")
for name, info in manifest["files"].items():
f = ckpt_dir / name
if not f.exists():
raise FileNotFoundError(f"checkpoint file missing: {name}")
actual = f.stat().st_size
if actual != info["bytes"]:
raise ValueError(f"{name}: size {actual} != manifest {info['bytes']}")
if sha256_file(f) != info["sha256"]:
raise ValueError(f"{name}: sha256 mismatch, checkpoint is corrupted")
config = CortexConfig.from_json(ckpt_dir / "config.json")
model = CortexForCausalLM(config)
weights = ckpt_dir / "weights.safetensors"
if weights.exists():
from safetensors.torch import load_file
state = load_file(str(weights))
else:
state = torch.load(ckpt_dir / "weights.pt", map_location="cpu", weights_only=True)
provenance = json.loads((ckpt_dir / "provenance.json").read_text(encoding="utf-8"))
# re-expand tied tensors: they were stored once and aliased, and must be written back
# into the state_dict under their original names before loading.
for alias, target in (provenance.get("tied_tensors") or {}).items():
if target not in state:
raise ValueError(f"checkpoint declares {alias} tied to missing tensor {target}")
state[alias] = state[target]
missing, unexpected = model.load_state_dict(state, strict=strict)
if strict and (missing or unexpected):
raise ValueError(f"state_dict mismatch: missing={missing} unexpected={unexpected}")
model.eval()
return model, config, provenance