File size: 4,771 Bytes
76d61a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/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()