File size: 1,853 Bytes
f1d3656 | 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 | #!/usr/bin/env python3
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
import torch
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.tiny_atmorep import TinyAtmoRep, TinyAtmoRepConfig
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--checkpoint", type=Path, default=ROOT / "weight" / "tiny_atmorep.pth")
parser.add_argument("--output", type=Path, default=ROOT / "result" / "prediction.pt")
parser.add_argument("--seed", type=int, default=17)
args = parser.parse_args()
payload = torch.load(args.checkpoint, map_location="cpu", weights_only=True)
config = TinyAtmoRepConfig(**payload["config"])
model = TinyAtmoRep(config)
model.load_state_dict(payload["model"])
model.eval()
torch.manual_seed(args.seed)
fields = torch.randn(1, *config.input_shape)
mask = torch.zeros(1, model.num_tokens, dtype=torch.bool)
mask[:, 1::4] = True
with torch.inference_mode():
ensemble = model(fields, mask, level=137.0)
target = model.tokenize(fields)
result = {
"ensemble": ensemble,
"ensemble_mean": ensemble.mean(dim=1),
"ensemble_std": ensemble.std(dim=1, unbiased=False),
"mask": mask,
"target": target,
"input_shape": tuple(fields.shape),
}
args.output.parent.mkdir(parents=True, exist_ok=True)
torch.save(target, args.output.parent / "target.pt")
torch.save(result, args.output)
print(json.dumps({
"output": str(args.output),
"ensemble_shape": list(ensemble.shape),
"mean_shape": list(result["ensemble_mean"].shape),
"finite": bool(torch.isfinite(ensemble).all()),
"bytes": args.output.stat().st_size,
}, indent=2))
if __name__ == "__main__":
main()
|