File size: 4,348 Bytes
3fef26f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
import argparse
from datetime import datetime, timedelta, timezone
import sys
from pathlib import Path

import numpy as np

PROJECT_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(PROJECT_DIR / "scripts"))
from config import load_config, project_path


VARIABLES = load_config()["data"]["variables"]


def prepare_input(dataset_dir, year, sample_index, onescience_src, input_steps, output_steps, model_height, model_width):
    import torch.nn.functional as F

    dataset_dir = Path(dataset_dir)
    required_paths = [
        dataset_dir / "data" / f"{year}.h5",
        dataset_dir / "stats" / "global_means.npy",
        dataset_dir / "stats" / "global_stds.npy",
    ]
    missing_paths = [path for path in required_paths if not path.is_file()]
    if missing_paths:
        missing = ", ".join(str(path) for path in missing_paths)
        raise FileNotFoundError(
            f"ERA5Dataset files are missing: {missing}. "
            "Run 'python scripts/fake_data.py' or update data.virtual_dir and "
            "inference.year in conf/config.yaml."
        )
    if onescience_src:
        sys.path.insert(0, onescience_src)
    from onescience.datapipes.climate.era5 import ERA5Dataset

    dataset = ERA5Dataset(
        dataset_dir=dataset_dir,
        used_years=[year],
        used_variables=VARIABLES,
        input_steps=input_steps,
        output_steps=output_steps,
        normalize=False,
    )
    fields, _, _, step_idx, _ = dataset[sample_index]
    if fields.ndim != 4 or tuple(fields.shape[:2]) != (input_steps, len(VARIABLES)):
        raise ValueError(f"Unexpected ERA5Dataset input shape: {tuple(fields.shape)}")
    fields = fields.float()
    fields[:, VARIABLES.index("tp")] = fields[:, VARIABLES.index("tp")].mul(1000.0).clamp(0.0, 1000.0)
    fields[:, VARIABLES.index("ttr")] /= 3600.0
    fields = F.interpolate(fields, size=(model_height, model_width), mode="bilinear", align_corners=False)
    return fields.unsqueeze(0).numpy().astype(np.float32), step_idx


def main():
    config = load_config()
    data_config = config["data"]
    model_config = config["model"]
    inference_config = config["inference"]
    parser = argparse.ArgumentParser(description="Run FuXi-S2S inference from OneScience ERA5Dataset data.")
    parser.add_argument("--model", default=str(project_path(model_config["path"])))
    parser.add_argument("--dataset-dir", default=str(project_path(data_config["virtual_dir"])))
    parser.add_argument("--year", type=int, default=inference_config["year"])
    parser.add_argument("--sample-index", type=int, default=inference_config["sample_index"])
    parser.add_argument("--onescience-src", help="Path containing the onescience Python package")
    parser.add_argument("--output-dir", default=str(project_path(inference_config["output_dir"])))
    parser.add_argument("--device", default=model_config["device"], choices=["cpu", "cuda", "dcu"])
    args = parser.parse_args()
    sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "model"))
    from fuxi_s2s import FuXiS2SModel

    model = FuXiS2SModel(args.model, device=args.device, providers=model_config["providers"])
    fields, step_idx = prepare_input(
        args.dataset_dir,
        args.year,
        args.sample_index,
        args.onescience_src,
        data_config["input_steps"],
        data_config["output_steps"],
        data_config["model_height"],
        data_config["model_width"],
    )
    if fields.shape[-2:] != (data_config["model_height"], data_config["model_width"]):
        raise ValueError(f"Unexpected FuXi-S2S input grid: {fields.shape[-2:]}")
    inputs = {"input": fields}
    if "step" in model.input_names:
        inputs["step"] = np.asarray([step_idx], dtype=np.float32)
    if "doy" in model.input_names:
        valid_time = datetime(args.year, 1, 1, tzinfo=timezone.utc) + timedelta(days=step_idx + 1)
        inputs["doy"] = np.asarray([min(365, valid_time.timetuple().tm_yday) / 365.0], dtype=np.float32)
    outputs = model(inputs)
    output_dir = Path(args.output_dir)
    output_dir.mkdir(parents=True, exist_ok=True)
    for name, value in outputs.items():
        np.save(output_dir / f"{name}.npy", value)
    print("Inference completed")
    print({name: tuple(value.shape) for name, value in outputs.items()})


if __name__ == "__main__":
    main()