CoolFace
Modelpublic

OneScience-Group/FourCastNet

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes40downloads
inference.py73 linesDownload Raw Back to scripts
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