#!/usr/bin/env python3 """Export a trusted training checkpoint into public, inference-only files.""" from __future__ import annotations import argparse import hashlib import json import os from pathlib import Path from typing import Any import torch import yaml from safetensors.torch import load_file, save_file def sha256(path: Path, chunk_size: int = 16 * 1024 * 1024) -> str: digest = hashlib.sha256() with path.open("rb") as stream: while chunk := stream.read(chunk_size): digest.update(chunk) return digest.hexdigest() def clean_state_dict(state: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: return { str(name): tensor.detach().cpu().contiguous() for name, tensor in state.items() } def tensor_summary(state: dict[str, torch.Tensor]) -> dict[str, int]: return { "parameter_tensors": len(state), "parameter_values": sum(int(tensor.numel()) for tensor in state.values()), } def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--checkpoint", required=True, type=Path) parser.add_argument("--release-config", required=True, type=Path) parser.add_argument("--output-dir", required=True, type=Path) parser.add_argument("--selected-epoch", type=int, default=121) parser.add_argument("--validation-threshold", type=float, default=0.1) args = parser.parse_args() # The source is a trusted, locally produced checkpoint. The public .pt file # written below contains only tensors and primitive metadata. source_payload = torch.load(args.checkpoint, map_location="cpu", weights_only=False) if not isinstance(source_payload, dict) or "model_state_dict" not in source_payload: raise ValueError("checkpoint must contain model_state_dict") source_epoch = int(source_payload.get("epoch", -1)) if source_epoch != args.selected_epoch: raise ValueError(f"checkpoint epoch {source_epoch} != selected epoch {args.selected_epoch}") with args.release_config.open("r", encoding="utf-8") as stream: release_config = yaml.safe_load(stream) if not isinstance(release_config, dict) or "model" not in release_config: raise ValueError("release config must contain a model section") state = clean_state_dict(source_payload["model_state_dict"]) args.output_dir.mkdir(parents=True, exist_ok=True) safe_path = args.output_dir / "model.safetensors" pt_path = args.output_dir / "best_model.pt" model_config_path = args.output_dir / "model_config.yaml" safe_tmp = args.output_dir / ".model.safetensors.partial" pt_tmp = args.output_dir / ".best_model.pt.partial" save_file( state, str(safe_tmp), metadata={"format": "pt", "model": "CI-Net", "epoch": str(source_epoch)}, ) os.replace(safe_tmp, safe_path) public_config: dict[str, Any] = { "model": release_config["model"], "input_sources": release_config.get("input_sources", ["concat", "hsr"]), "input_window": release_config.get("input_window", {"past_minutes": 50, "interval_minutes": 10}), } public_payload = { "epoch": source_epoch, "model_state_dict": state, "config": public_config, } torch.save(public_payload, pt_tmp) os.replace(pt_tmp, pt_path) with model_config_path.open("w", encoding="utf-8") as stream: yaml.safe_dump(public_config, stream, sort_keys=False) safe_state = load_file(str(safe_path), device="cpu") if state.keys() != safe_state.keys(): raise RuntimeError("safetensors key mismatch") for name in state: if not torch.equal(state[name], safe_state[name]): raise RuntimeError(f"safetensors tensor mismatch: {name}") metadata = { "model": "CI-Net SimVP with BT auxiliary prediction", "selected_epoch": source_epoch, "selection_dataset": "2024 validation period", "validation_threshold": float(args.validation_threshold), "configuration": public_config, "source_checkpoint": { "filename": args.checkpoint.name, "sha256": sha256(args.checkpoint), }, "release_files": { safe_path.name: {"sha256": sha256(safe_path), "bytes": safe_path.stat().st_size}, pt_path.name: {"sha256": sha256(pt_path), "bytes": pt_path.stat().st_size}, }, **tensor_summary(state), "tensor_equality_verified": True, } with (args.output_dir / "checkpoint_metadata.json").open("w", encoding="utf-8") as stream: json.dump(metadata, stream, indent=2, ensure_ascii=False) stream.write("\n") print(json.dumps(metadata, indent=2)) if __name__ == "__main__": main()