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