Download scripts/inference.py from OneScience-Group/FuXi_v21: direct link, hf CLI and curl.
- Browser
- Download file 5.98 kB
-
https://huggingface.co/OneScience-Group/FuXi_v21/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/FuXi_v21/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/FuXi_v21/resolve/main/scripts/inference.py
5.98 kB
| """Autoregressive inference using a project-produced FuXi 2.1 checkpoint.""" | |
| from __future__ import annotations | |
| import argparse | |
| from datetime import datetime, timedelta | |
| import numpy as np | |
| import torch | |
| import xarray as xr | |
| from onescience.datapipes.climate.era5 import ERA5Dataset | |
| from common import load_config, resolve_path | |
| from model.FuXi21 import FuXi21 | |
| from variables import c85_from_config | |
| CHECKPOINT_FORMAT = "fuxi21_reconstructed_checkpoint_v1" | |
| def select_device(requested: str) -> torch.device: | |
| if requested not in {"auto", "cpu", "cuda"}: | |
| raise ValueError("inference.device must be auto, cpu, or cuda") | |
| if requested == "cuda" or (requested == "auto" and torch.cuda.is_available()): | |
| if not torch.cuda.is_available(): | |
| raise RuntimeError("inference.device=cuda, but no CUDA/HIP device is available") | |
| return torch.device("cuda") | |
| return torch.device("cpu") | |
| def load_array(path_value: str | None, cfg: dict, shape: tuple[int, ...], name: str) -> torch.Tensor: | |
| if path_value is None: | |
| raise ValueError(f"model.{name}_file is required outside the smoke profile") | |
| path = resolve_path(path_value, cfg) | |
| value = torch.from_numpy(np.load(path)).float() | |
| if tuple(value.shape) != shape: | |
| raise ValueError(f"{name} must have shape {shape}, got {tuple(value.shape)}") | |
| return value | |
| def build_model(cfg: dict) -> FuXi21: | |
| model_cfg = cfg["model"] | |
| profile_name = model_cfg["profile"] | |
| profile = model_cfg["profiles"][profile_name] | |
| height, width = profile["grid_size"] | |
| if profile_name == "smoke": | |
| static_fields = torch.zeros(6, height, width) | |
| channel_mask = torch.ones(85, height, width) | |
| else: | |
| static_fields = load_array(model_cfg["static_fields_file"], cfg, (6, height, width), "static_fields") | |
| channel_mask = load_array(model_cfg["channel_mask_file"], cfg, (85, height, width), "channel_mask") | |
| return FuXi21( | |
| static_fields, | |
| channel_mask, | |
| activation_checkpointing=False, | |
| **profile, | |
| ) | |
| def temporal_features(valid_time: datetime, step: int, device: torch.device) -> tuple[torch.Tensor, ...]: | |
| return ( | |
| torch.tensor([step], device=device, dtype=torch.float32), | |
| torch.tensor([(valid_time.hour * 60 + valid_time.minute) / 1440], device=device), | |
| torch.tensor([min(365, valid_time.timetuple().tm_yday) / 365], device=device), | |
| ) | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--config", default="conf/config.yaml") | |
| parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default=None) | |
| parser.add_argument("--preflight-only", action="store_true") | |
| args = parser.parse_args() | |
| cfg = load_config(args.config) | |
| if cfg.get("protocol") != "non_official_protocol": | |
| raise ValueError("Inference config must declare protocol: non_official_protocol") | |
| infer_cfg = cfg["inference"] | |
| channels, diagnostics = c85_from_config(cfg) | |
| checkpoint_path = resolve_path(infer_cfg["checkpoint"], cfg) | |
| if not checkpoint_path.is_file(): | |
| raise FileNotFoundError(f"Project checkpoint not found: {checkpoint_path}") | |
| device = select_device(args.device or infer_cfg["device"]) | |
| checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=True) | |
| if checkpoint.get("format") != CHECKPOINT_FORMAT: | |
| raise ValueError(f"Checkpoint must use format {CHECKPOINT_FORMAT}") | |
| if checkpoint.get("protocol") != "non_official_protocol": | |
| raise ValueError("Checkpoint protocol must be non_official_protocol") | |
| if checkpoint.get("model_profile") != cfg["model"]["profile"]: | |
| raise ValueError("Checkpoint model profile does not match the configured model profile") | |
| if args.preflight_only: | |
| print(f"checkpoint={checkpoint_path}, profile={checkpoint['model_profile']}, device={device}") | |
| return | |
| model = build_model(cfg).to(device) | |
| model.load_state_dict(checkpoint["model"]) | |
| model.eval() | |
| split = infer_cfg["split"] | |
| split_cfg = cfg["data"]["splits"][split] | |
| dataset = ERA5Dataset( | |
| dataset_dir=str(resolve_path(cfg["paths"]["data_root"], cfg)), | |
| used_years=split_cfg["years"], | |
| used_variables=channels, | |
| input_steps=cfg["data"]["input_steps"], | |
| output_steps=cfg["data"]["output_steps"], | |
| normalize=True, | |
| ) | |
| state, _, _, _, time_index = dataset[0] | |
| crop_size = cfg["data"]["crop_size"] | |
| if crop_size is not None: | |
| state = state[..., : crop_size[0], : crop_size[1]] | |
| state = state.unsqueeze(0).to(device) | |
| valid_time = datetime.strptime(time_index[-1], "%Y%m%d%H") | |
| interval = timedelta(hours=cfg["data"]["time_step_hours"]) | |
| diagnostic_indices = [channels.index(name) for name in diagnostics] | |
| forecasts = [] | |
| valid_times = [] | |
| for step in range(infer_cfg["steps"]): | |
| with torch.inference_mode(): | |
| state = model(state, *temporal_features(valid_time, step, device)) | |
| forecasts.append(state[:, -1].float().cpu().numpy()[0]) | |
| valid_times.append(np.datetime64(valid_time)) | |
| if infer_cfg["zero_diagnostic_feedback"]: | |
| state[:, -1, diagnostic_indices] = 0 | |
| valid_time += interval | |
| height, width = forecasts[0].shape[-2:] | |
| output_path = resolve_path(infer_cfg["output_file"], cfg) | |
| output_path.parent.mkdir(parents=True, exist_ok=True) | |
| xr.DataArray( | |
| np.stack(forecasts), | |
| dims=("time", "channel", "lat", "lon"), | |
| coords={ | |
| "time": valid_times, | |
| "channel": channels, | |
| "lat": np.linspace(90, -90, height), | |
| "lon": np.arange(width) * (360 / width), | |
| }, | |
| attrs={"checkpoint_format": CHECKPOINT_FORMAT, "protocol": cfg["protocol"]}, | |
| name="forecast", | |
| ).to_netcdf(output_path) | |
| print(f"Saved {len(forecasts)} forecast step(s) to {output_path}") | |
| if __name__ == "__main__": | |
| main() | |