EchoCast-3D / scripts /inference.py
zhangrenchao's picture
Publish EchoCast-3D reproduction
e0a6aa0 verified
Raw
History Blame Contribute Delete
1.98 kB
"""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()