"""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:: / 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