OneScience-Group/FourCastNet
040
1import sys2from pathlib import Path3 4# 获取项目根目录(train.py上级的上级)5root_path = Path(__file__).parent.parent6sys.path.append(str(root_path))7import torch8import os9import glob10import numpy as np11import h5py12from tqdm import tqdm13from model.fourcastnet import FourCastNet14from onescience.utils.YParams import YParams15from onescience.datapipes.climate import ERA5Datapipe16 17 18def get_stats(data_dir, channels):19 """从新版 h5 中读取变量列表与归一化参数(均值/标准差)"""20 h5_files = sorted(glob.glob(os.path.join(data_dir, "data", "*.h5")))21 with h5py.File(h5_files[0], "r") as f:22 ds = f["fields"]23 all_variables = [v.decode() if isinstance(v, bytes) else v for v in ds.attrs["variables"]]24 mu = f["global_means"][:] # [1, C, 1, 1]25 std = f["global_stds"][:]26 27 channel_indices = [all_variables.index(v) for v in channels]28 means = mu[:, channel_indices, :, :]29 stds = std[:, channel_indices, :, :]30 return means, stds31 32 33if __name__ == "__main__":34 current_path = os.getcwd()35 sys.path.append(current_path)36 37 ## Model config init38 config_file_path = os.path.join(current_path, "conf/config.yaml")39 cfg = YParams(config_file_path, "model")40 41 ## DataLoader init42 cfg_data = YParams(config_file_path, "datapipe")43 means, stds = get_stats(cfg_data.dataset.data_dir, cfg_data.dataset.channels)44 45 cfg['N_in_channels'] = len(cfg_data.dataset.channels)46 cfg['N_out_channels'] = len(cfg_data.dataset.channels)47 48 datapipe = ERA5Datapipe(49 dataset_dir=cfg_data.dataset.data_dir,50 used_variables=cfg_data.dataset.channels,51 used_years=cfg_data.dataset.test_time,52 distributed=False,53 batch_size=1,54 num_workers=4,55 )56 test_dataloader, _ = datapipe.get_dataloader("test")57 58 ckpt = torch.load(f"{cfg.checkpoint_dir}/model_bak.pth", map_location="cuda:0")59 model = FourCastNet().to('cuda:0')60 model.load_state_dict(ckpt["model_state_dict"])61 62 model.eval()63 os.makedirs('result/output/', exist_ok=True)64 print(f"📂 infer results will be generated to './result/output/'")65 with torch.no_grad():66 for data in tqdm(test_dataloader, desc="Inferring testset", unit="batch"):67 invar = data[0].to('cuda:0', dtype=torch.float32)68 filename = data[4][-1][0]69 invar = invar[:, :, :-1, :]70 pred_var = model(invar).cpu().numpy()71 pred_var = pred_var * stds + means72 np.save(f"result/output/{filename}.npy", pred_var)73 