Download scripts/inference.py from OneScience-Group/FuXi-S2S: direct link, hf CLI and curl.
- Browser
- Download file 4.35 kB
-
https://huggingface.co/OneScience-Group/FuXi-S2S/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/FuXi-S2S/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/FuXi-S2S/resolve/main/scripts/inference.py
4.35 kB
| 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() | |