CoolFace
Modelpublic

OneScience-Group/Pangu_Weather

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes11downloads
inference.py91 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 sys10import glob11import numpy as np12import h5py13from tqdm import tqdm14from model.pangu import Pangu15from onescience.utils.YParams import YParams16from onescience.datapipes.climate import ERA5Datapipe17 18 19def get_stats(data_dir, channels):20    """从新版 h5 中读取变量列表与归一化参数(均值/标准差)"""21    h5_files = sorted(glob.glob(os.path.join(data_dir, "data", "*.h5")))22    with h5py.File(h5_files[0], "r") as f:23        ds = f["fields"]24        all_variables = [v.decode() if isinstance(v, bytes) else v for v in ds.attrs["variables"]]25        mu = f["global_means"][:]   # [1, C, 1, 1]26        std = f["global_stds"][:]27 28    channel_indices = [all_variables.index(v) for v in channels]29    means = mu[:, channel_indices, :, :]30    stds = std[:, channel_indices, :, :]31    return means, stds32 33 34if __name__ == "__main__":35    current_path = os.getcwd()36    sys.path.append(current_path)37 38    ## Model config init39    config_file_path = os.path.join(current_path, "conf/config.yaml")40    cfg = YParams(config_file_path, "model")41    ## DataLoader init42    cfg_data = YParams(config_file_path, "datapipe")43 44    means, stds = get_stats(cfg_data.dataset.data_dir, cfg_data.dataset.channels)45 46    datapipe = ERA5Datapipe(47        dataset_dir=cfg_data.dataset.data_dir,48        used_variables=cfg_data.dataset.channels,49        used_years=cfg_data.dataset.test_time,50        distributed=False,51        batch_size=1,52        num_workers=4,53    )54    test_dataloader, _ = datapipe.get_dataloader("test")55 56    static_dir = os.path.join(cfg_data.dataset.data_dir, "static")57    58    land_mask = torch.from_numpy(np.load(os.path.join(static_dir, "land_mask.npy")).astype(np.float32))59    soil_type = torch.from_numpy(np.load(os.path.join(static_dir, "soil_type.npy")).astype(np.float32))60    topography = torch.from_numpy(np.load(os.path.join(static_dir, "topography.npy")).astype(np.float32))61    topography = (topography - topography.mean()) / (topography.std(unbiased=False) + 1e-6)62    surface_mask = torch.stack([land_mask, soil_type, topography], dim=0).to('cuda:0')63    surface_mask = surface_mask.unsqueeze(0).repeat(cfg_data.dataloader.batch_size, 1, 1, 1)64 65    ckpt = torch.load(f"{cfg.checkpoint_dir}/model_bak.pth", map_location="cuda:0")66    model = Pangu(img_size=cfg_data.dataset.img_size,67                  patch_size=cfg.patch_size,68                  embed_dim=cfg.embed_dim,69                  num_heads=cfg.num_heads,70                  window_size=cfg.window_size,71                  ).to('cuda:0')72    model.load_state_dict(ckpt["model_state_dict"])73 74    model.eval()75    os.makedirs('result/output/', exist_ok=True)76    print(f"📂 samples will be generated to './result/output/'")77    with torch.no_grad():78        for data in tqdm(test_dataloader, desc="Inferring testset", unit="batch"):79            invar = data[0]80            outvar = data[1]81            filename = data[4][-1][0]82            invar_surface = invar[:, :4, :, :].to("cuda:0", dtype=torch.float32)83            invar_upper_air = invar[:, 4:, :, :].to("cuda:0", dtype=torch.float32)84            invar = torch.concat([invar_surface, surface_mask, invar_upper_air], dim=1)85 86            out_surface, out_upper_air = model(invar)87            out_upper_air = out_upper_air.reshape(invar_upper_air.shape)88            pred_var = torch.concat([out_surface, out_upper_air], dim=1).cpu().numpy()89            pred_var = pred_var * stds + means90            np.save(f"result/output/{filename}.npy", pred_var)91