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