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