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