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