File size: 1,982 Bytes
e0a6aa0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Generate native packed-wedge ensemble forecasts."""

from __future__ import annotations

import argparse
from pathlib import Path

import numpy as np
import torch
import sys
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))

from model.echocast_3d import PackedRadarDataset, load_config, DiffusionSchedule, EchoCast3D


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", default="conf/config.yaml")
    parser.add_argument("--data", default="data")
    parser.add_argument("--checkpoint", default="result/checkpoints/echocast_3d.pt")
    parser.add_argument("--output", default="result/output/predictions.npz")
    args = parser.parse_args()
    config = load_config(ROOT / args.config)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model = EchoCast3D(config).to(device)
    checkpoint = torch.load(ROOT / args.checkpoint, map_location=device, weights_only=False)
    model.load_state_dict(checkpoint["model"])
    model.eval()
    sample = PackedRadarDataset(ROOT / args.data, "validation")[0]
    history = sample["values"][None, :3].float().to(device)
    observed = sample["observed"][None, :3].bool().to(device)
    schedule = DiffusionSchedule(device=device, **{k: config["diffusion"][k] for k in ("steps", "beta_start", "beta_end")})
    ensemble = []
    for member in range(config["diffusion"]["ensemble_size"]):
        torch.manual_seed(config["seed"] + member)
        ensemble.append(schedule.sample(model, history, observed)[0].cpu().numpy())
    output = ROOT / args.output
    output.parent.mkdir(parents=True, exist_ok=True)
    np.savez(
        output,
        ensemble=np.stack(ensemble),
        truth=sample["truth"][3:].numpy(),
        validity=sample["validity"][3:].numpy(),
        observed_history=sample["observed"][:3].numpy(),
    )
    print(f"saved {output}: ensemble={len(ensemble)}, steps={schedule.steps}")


if __name__ == "__main__":
    main()