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