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