File size: 3,688 Bytes
9333192 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 | import sys
from pathlib import Path
# 获取项目根目录(train.py上级的上级)
root_path = Path(__file__).parent.parent
sys.path.append(str(root_path))
import torch
import os
import sys
import glob
import numpy as np
import h5py
from tqdm import tqdm
from model.pangu import Pangu
from onescience.utils.YParams import YParams
from onescience.datapipes.climate import ERA5Datapipe
def get_stats(data_dir, channels):
"""从新版 h5 中读取变量列表与归一化参数(均值/标准差)"""
h5_files = sorted(glob.glob(os.path.join(data_dir, "data", "*.h5")))
with h5py.File(h5_files[0], "r") as f:
ds = f["fields"]
all_variables = [v.decode() if isinstance(v, bytes) else v for v in ds.attrs["variables"]]
mu = f["global_means"][:] # [1, C, 1, 1]
std = f["global_stds"][:]
channel_indices = [all_variables.index(v) for v in channels]
means = mu[:, channel_indices, :, :]
stds = std[:, channel_indices, :, :]
return means, stds
if __name__ == "__main__":
current_path = os.getcwd()
sys.path.append(current_path)
## Model config init
config_file_path = os.path.join(current_path, "conf/config.yaml")
cfg = YParams(config_file_path, "model")
## DataLoader init
cfg_data = YParams(config_file_path, "datapipe")
means, stds = get_stats(cfg_data.dataset.data_dir, cfg_data.dataset.channels)
datapipe = ERA5Datapipe(
dataset_dir=cfg_data.dataset.data_dir,
used_variables=cfg_data.dataset.channels,
used_years=cfg_data.dataset.test_time,
distributed=False,
batch_size=1,
num_workers=4,
)
test_dataloader, _ = datapipe.get_dataloader("test")
static_dir = os.path.join(cfg_data.dataset.data_dir, "static")
land_mask = torch.from_numpy(np.load(os.path.join(static_dir, "land_mask.npy")).astype(np.float32))
soil_type = torch.from_numpy(np.load(os.path.join(static_dir, "soil_type.npy")).astype(np.float32))
topography = torch.from_numpy(np.load(os.path.join(static_dir, "topography.npy")).astype(np.float32))
topography = (topography - topography.mean()) / (topography.std(unbiased=False) + 1e-6)
surface_mask = torch.stack([land_mask, soil_type, topography], dim=0).to('cuda:0')
surface_mask = surface_mask.unsqueeze(0).repeat(cfg_data.dataloader.batch_size, 1, 1, 1)
ckpt = torch.load(f"{cfg.checkpoint_dir}/model_bak.pth", map_location="cuda:0")
model = Pangu(img_size=cfg_data.dataset.img_size,
patch_size=cfg.patch_size,
embed_dim=cfg.embed_dim,
num_heads=cfg.num_heads,
window_size=cfg.window_size,
).to('cuda:0')
model.load_state_dict(ckpt["model_state_dict"])
model.eval()
os.makedirs('result/output/', exist_ok=True)
print(f"📂 samples will be generated to './result/output/'")
with torch.no_grad():
for data in tqdm(test_dataloader, desc="Inferring testset", unit="batch"):
invar = data[0]
outvar = data[1]
filename = data[4][-1][0]
invar_surface = invar[:, :4, :, :].to("cuda:0", dtype=torch.float32)
invar_upper_air = invar[:, 4:, :, :].to("cuda:0", dtype=torch.float32)
invar = torch.concat([invar_surface, surface_mask, invar_upper_air], dim=1)
out_surface, out_upper_air = model(invar)
out_upper_air = out_upper_air.reshape(invar_upper_air.shape)
pred_var = torch.concat([out_surface, out_upper_air], dim=1).cpu().numpy()
pred_var = pred_var * stds + means
np.save(f"result/output/{filename}.npy", pred_var)
|