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