File size: 5,578 Bytes
954544e
 
 
 
 
 
 
 
2cb6966
954544e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2cb6966
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import hashlib
import json
import os
from pathlib import Path
from typing import Any

from safetensors import safe_open
from safetensors.torch import load_file, save_file
import torch

from .model import DotRecurrentDepthModel


def _sha256(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as handle:
        for chunk in iter(lambda: handle.read(8 * 1024 * 1024), b""):
            digest.update(chunk)
    return digest.hexdigest()


def export_reasoning_core(checkpoint_path: str | Path, output_dir: str | Path) -> dict[str, Any]:
    checkpoint = Path(checkpoint_path)
    payload = torch.load(checkpoint, map_location="cpu", weights_only=False)
    if payload.get("format_version") != 1:
        raise ValueError(f"unsupported checkpoint format: {payload.get('format_version')}")
    state = payload.get("reasoning_core")
    if not isinstance(state, dict) or not state:
        raise ValueError("checkpoint has no reasoning_core state")
    tensors = {name: tensor.detach().contiguous().cpu() for name, tensor in state.items()}

    root = Path(output_dir)
    root.mkdir(parents=True, exist_ok=True)
    final_weights = root / "reasoning_core.safetensors"
    temporary_weights = root / "reasoning_core.safetensors.tmp"
    save_file(tensors, temporary_weights)
    os.replace(temporary_weights, final_weights)
    manifest = {
        "format_version": 1,
        "artifact": "Dot recurrent-depth reasoning core",
        "source_checkpoint": str(checkpoint),
        "source_step": int(payload["step"]),
        "weights": final_weights.name,
        "weights_sha256": _sha256(final_weights),
        "runtime": {
            "base_model": ".",
            "loader": "dot_rd.export.load_exported_core",
            "use_cache": False,
        },
        "architecture": payload["architecture"],
        "training_metrics": payload.get("metrics") or {},
    }
    final_manifest = root / "dot_recurrent_manifest.json"
    temporary_manifest = root / "dot_recurrent_manifest.json.tmp"
    temporary_manifest.write_text(json.dumps(manifest, indent=2) + "\n", encoding="utf-8")
    os.replace(temporary_manifest, final_manifest)
    return manifest


def load_exported_core(path: str | Path, model: DotRecurrentDepthModel) -> dict[str, Any]:
    root = Path(path)
    manifest = json.loads((root / "dot_recurrent_manifest.json").read_text(encoding="utf-8"))
    weights_path = root / manifest["weights"]
    actual_hash = _sha256(weights_path)
    if actual_hash != manifest["weights_sha256"]:
        raise ValueError(
            f"reasoning core hash mismatch: expected {manifest['weights_sha256']}, got {actual_hash}"
        )
    expected = model.architecture_manifest()
    actual = manifest.get("architecture") or {}
    for field in ("architecture", "base_parameter_count", "total_parameter_count", "new_parameter_count"):
        if actual.get(field) != expected.get(field):
            raise ValueError(f"exported core architecture mismatch for {field}")
    model.reasoning_core.load_state_dict(load_file(weights_path, device="cpu"), strict=True)
    return manifest


def load_inference_checkpoint(
    path: str | Path, model: DotRecurrentDepthModel
) -> dict[str, Any]:
    """Load a verified Dot core plus its explicitly saved backbone repair delta."""

    root = Path(path)
    manifest = json.loads((root / "manifest.json").read_text(encoding="utf-8"))
    weights_path = root / manifest["weights"]
    actual_hash = _sha256(weights_path)
    if actual_hash != manifest["weights_sha256"]:
        raise ValueError(
            f"inference checkpoint hash mismatch: expected {manifest['weights_sha256']}, "
            f"got {actual_hash}"
        )
    expected = model.architecture_manifest()
    actual = manifest.get("architecture") or {}
    for field in ("architecture", "base_parameter_count", "total_parameter_count", "new_parameter_count"):
        if actual.get(field) != expected.get(field):
            raise ValueError(f"inference checkpoint architecture mismatch for {field}")

    core_state = model.reasoning_core.state_dict()
    backbone_state = model.backbone.state_dict()
    seen_core: set[str] = set()
    seen_backbone: set[str] = set()
    with safe_open(weights_path, framework="pt", device="cpu") as handle:
        for key in handle.keys():
            if key.startswith("reasoning_core."):
                name = key.removeprefix("reasoning_core.")
                target = core_state.get(name)
                seen_core.add(name)
            elif key.startswith("backbone_delta."):
                name = key.removeprefix("backbone_delta.")
                target = backbone_state.get(name)
                seen_backbone.add(name)
            else:
                raise ValueError(f"unsupported inference checkpoint tensor: {key}")
            if target is None:
                raise ValueError(f"checkpoint tensor does not exist in Dot: {key}")
            tensor = handle.get_tensor(key)
            if tensor.shape != target.shape:
                raise ValueError(f"checkpoint shape mismatch for {key}: {tensor.shape} != {target.shape}")
            with torch.no_grad():
                target.copy_(tensor.to(device=target.device, dtype=target.dtype))
    missing_core = sorted(set(core_state).difference(seen_core))
    if missing_core:
        raise ValueError(f"inference checkpoint is missing core tensors: {missing_core[:5]}")
    if not seen_backbone:
        raise ValueError("inference checkpoint has no backbone repair delta")
    return manifest