Download code/training/src/export_weights.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 4.77 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/export_weights.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/export_weights.py
-
curl -L -o export_weights.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/export_weights.py
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() | |