SEEDS / scripts /inference.py
yzt15806542928's picture
Upload folder using huggingface_hub
1ca0208 verified
Raw
History Blame Contribute Delete
2.83 kB
"""Generate conditional SEEDS ensemble forecasts."""
from __future__ import annotations
import argparse
from pathlib import Path
import numpy as np
import torch
from common import build_model, choose_device, load_config, resolve_path, set_seed
from data_loader import build_dataloader
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config", default="conf/config.yaml")
parser.add_argument("--checkpoint", default=None)
parser.add_argument("--members", type=int, default=None)
parser.add_argument("--member-batch-size", type=int, default=None)
parser.add_argument("--steps", type=int, default=None)
parser.add_argument("--max-batches", type=int, default=None)
parser.add_argument("--device", default=None)
args = parser.parse_args()
config = load_config(args.config)
set_seed(config["project"]["seed"])
device = choose_device(args.device or config["training"]["device"])
data, paths, training = config["data"], config["paths"], config["training"]
loader = build_dataloader(resolve_path(paths["test_data"], args.config), len(data["variables"]), data["faces"], data["height"], data["width"], data["seed_count"], config["sampling"]["data_batch_size"], False, training["num_workers"], args.max_batches)
model = build_model(config).to(device)
checkpoint = resolve_path(args.checkpoint or paths["checkpoint"], args.config)
state = torch.load(checkpoint, map_location=device, weights_only=False)
model.load_state_dict(state["model"] if "model" in state else state)
model.eval()
members = args.members or data["target_count"]
steps = args.steps or config["sampling"]["steps"]
member_batch_size = args.member_batch_size or config["sampling"]["member_batch_size"]
outputs, targets, seed_values = [], [], []
with torch.no_grad():
for batch in loader:
seeds = batch["seeds"].to(device)
climate = batch["climate"].to(device)
outputs.append(model.sample(seeds, climate, members=members, steps=steps, member_batch_size=member_batch_size).cpu().numpy())
targets.append(batch["targets"].numpy())
seed_values.append(batch["seeds"].numpy())
if not outputs:
raise RuntimeError("test loader produced no batches")
output_dir = resolve_path(paths["result_dir"], args.config) / "output"
output_dir.mkdir(parents=True, exist_ok=True)
prediction = np.concatenate(outputs, axis=0)
target = np.concatenate(targets, axis=0)
np.save(output_dir / "prediction.npy", prediction)
np.save(output_dir / "target.npy", target)
np.save(output_dir / "seeds.npy", np.concatenate(seed_values, axis=0))
print(f"saved prediction shape={prediction.shape} to {output_dir}")
if __name__ == "__main__":
main()