ci-net / code /training /src /export_weights.py
lsh9034's picture
Add files using upload-large-folder tool
76d61a0 verified
Raw History Blame Contribute Delete
4.77 kB
#!/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()