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()
|