File size: 2,831 Bytes
1ca0208
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
52
53
54
55
56
57
58
59
60
"""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()