File size: 7,542 Bytes
6fbe100
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
"""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