Team Ai
Modelpublic

OneScience-Group/Pangu_Weather

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes38downloads
result.py227 linesDownload Raw Back to scripts
1import numpy as np2import matplotlib.pyplot as plt3import os4import sys5import glob6import h5py7from datetime import datetime8from tqdm import tqdm9from onescience.utils.fcn.YParams import YParams10from matplotlib import rcParams11 12# rcParams['font.family'] = 'serif'13# rcParams['font.serif'] = ['DejaVu Serif']14rcParams['mathtext.fontset'] = 'stix'15rcParams['axes.linewidth'] = 0.916rcParams['xtick.major.width'] = 0.917rcParams['ytick.major.width'] = 0.918 19 20def get_metadata(data_dir, channels):21    """从新版 h5 attrs 中读取变量列表和 time_step"""22    h5_files = sorted(glob.glob(os.path.join(data_dir, "data", "*.h5")))23    with h5py.File(h5_files[0], "r") as f:24        ds = f["fields"]25        all_variables = [v.decode() if isinstance(v, bytes) else v for v in ds.attrs["variables"]]26        time_step = int(ds.attrs["time_step"])27 28    channel_indices = [all_variables.index(v) for v in channels]29 30    total_files = [f for f in os.listdir('./result/output/') if f.endswith('.npy')]31    total_files.sort()32    return total_files, channel_indices, time_step33 34 35def filename_to_index(filename, time_step):36    """将 YYYYMMDDHH 格式的文件名转换为年度 h5 文件中的时间步索引"""37    dt = datetime.strptime(filename, "%Y%m%d%H")38    year_start = datetime(dt.year, 1, 1)39    hours = (dt - year_start).total_seconds() / 360040    return int(hours / time_step)41 42 43def group_files_by_year(total_files, time_step):44    """将输出文件按年份归组,减少重复打开年度 h5 文件的开销"""45    files_by_year = {}46    for file in total_files:47        fname = file[:-4]48        year = fname[:4]49        files_by_year.setdefault(year, []).append((file, filename_to_index(fname, time_step)))50    return files_by_year51 52 53def get_result(total_files, channel_indices, time_step, data_dir, clim_mean):54    channel_rmse = np.zeros(len(channel_indices))55    channel_acc = np.zeros(len(channel_indices))56    clim_mean = clim_mean[0, :, :, :]57    if not os.path.exists('./result/rmse.npy') or not os.path.exists('result/acc.npy'):58        numerator = np.zeros(len(channel_indices))59        pred_sq_sum = np.zeros(len(channel_indices))60        label_sq_sum = np.zeros(len(channel_indices))61        files_by_year = group_files_by_year(total_files, time_step)62        with tqdm(total=len(total_files), unit="files") as pbar:63            for year, year_files in files_by_year.items():64                with h5py.File(os.path.join(data_dir, 'data', f'{year}.h5'), "r") as f:65                    fields = f["fields"]66                    for file, t_idx in year_files:67                        label = fields[t_idx]  # [C, H, W]68                        label = label[channel_indices]69                        pred = np.load(f'result/output/{file}').squeeze()70 71                        label_anom = label - clim_mean72                        pred_anom = pred - clim_mean73                        # 累加74                        numerator += np.sum(pred_anom * label_anom, axis=(1, 2))75                        pred_sq_sum += np.sum(pred_anom ** 2, axis=(1, 2))76                        label_sq_sum += np.sum(label_anom ** 2, axis=(1, 2))77 78                        channel_rmse += np.sqrt(np.mean((label - pred) ** 2, axis=(1, 2)))79                        pbar.update(1)80        channel_rmse /= len(total_files)81        channel_acc = numerator / (np.sqrt(pred_sq_sum * label_sq_sum) + 1e-8)82        np.save('./result/acc.npy', channel_acc)83        np.save('./result/rmse.npy', channel_rmse)84 85 86def show_result():87    channel_rmse = np.load('./result/rmse.npy')88    channel_acc = np.load('./result/acc.npy')89 90    channels = [cfg_data.dataset.channels[i] for i in range(len(channel_indices))]91    w = 24  # 最长 channel 名宽度92 93    # 表头94    print(f"┌{'─' * (w + 2)}┬{'─' * 14}┬{'─' * 14}┐")95    print(f"│ {'Channel':<{w}} │ {'RMSE':>12} │ {'ACC':>12} │")96    print(f"├{'─' * (w + 2)}┼{'─' * 14}┼{'─' * 14}┤")97    # 数据行98    for i, ch in enumerate(channels):99        print(f"│ {ch:<{w}} │ {channel_rmse[i]:>12.4f} | {channel_acc[i]:>12.4f} |")100    print(f"├{'─' * (w + 2)}┼{'─' * 14}┼{'─' * 14}┤")101    print(f"│ {'Average':<{w}} │ {np.mean(channel_rmse):>12.4f} │ {np.mean(channel_acc):>12.4f} │")102    print(f"└{'─' * (w + 2)}┴{'─' * 14}┴{'─' * 14}┘")103 104 105def plot(label, pred, var, filename):106    # 基础设置107    fig, axes = plt.subplots(1, 3, figsize=(15, 4))108 109    # 坐标轴标签110    xtick_labels = ['180°W', '90°W', '0°', '90°E', '180°E']111    ytick_labels = ['90°S', '45°S', '0°', '45°N', '90°N']112    xticks = np.linspace(0, label.shape[-1] - 1, 5)113    yticks = np.linspace(0, label.shape[-2] - 1, 5)114 115    # 计算统一色条范围116    vmin = min(label.min(), pred.min())117    vmax = max(label.max(), pred.max())118 119    # 计算差异和 RMSE120    diff = label - pred121    rmse = np.sqrt(np.mean(diff ** 2))122    diff_abs_max = np.abs(diff).max()123 124    # 绘图配置125    plot_configs = [126        {'data': label, 'title': 'Truth', 'cmap': 'viridis', 'vmin': vmin, 'vmax': vmax},127        {'data': pred,  'title': 'Prediction', 'cmap': 'viridis', 'vmin': vmin, 'vmax': vmax},128        {'data': diff,  'title': f'Difference (RMSE={rmse:.2f})', 'cmap': 'RdBu_r', 'vmin': -diff_abs_max, 'vmax': diff_abs_max},129    ]130 131    # 统一绘制132    for ax, cfg in zip(axes, plot_configs):133        im = ax.imshow(cfg['data'], cmap=cfg['cmap'], vmin=cfg['vmin'], vmax=cfg['vmax'])134        ax.set_title(cfg['title'], fontsize=12, pad=4)135        ax.set_xlabel('Longitude')136        ax.set_ylabel('Latitude')137        ax.set_xticks(xticks)138        ax.set_xticklabels(xtick_labels)139        ax.set_yticks(yticks)140        ax.set_yticklabels(ytick_labels)141        plt.colorbar(im, ax=ax, orientation='horizontal')142 143    # 总标题144    fig.suptitle(var, fontsize=14, fontweight='bold', y=0.98)145 146    plt.savefig(filename, dpi=300, bbox_inches='tight')147    plt.close()148 149 150def plot_loss(train_loss, valid_loss):151 152    mask = ~(np.isnan(train_loss) | np.isnan(valid_loss))153    train_loss = train_loss[mask]154    valid_loss = valid_loss[mask]155 156    fig, ax = plt.subplots(figsize=(5, 3.5))157    # 配置158    colors = {'train': '#2563EB', 'valid': '#EA580C'}159    epochs = np.arange(1, len(train_loss) + 1)160 161    # 绑定曲线162    ax.plot(epochs, train_loss, color=colors['train'], linewidth=1.5, label='Train')163    ax.plot(epochs, valid_loss, color=colors['valid'], linewidth=1.5, label='Valid', linestyle='--')164    # 标注最小值165    min_idx = np.argmin(valid_loss)166    ax.scatter(epochs[min_idx], valid_loss[min_idx],167               color=colors['valid'], s=40, zorder=5, edgecolors='white')168    ax.annotate(f'Best: {valid_loss[min_idx]:.3f}',169                xy=(epochs[min_idx], valid_loss[min_idx]),170                xytext=(10, 10), textcoords='offset points', fontsize=8, color=colors['valid'],171                arrowprops=dict(arrowstyle='-', color=colors['valid'], lw=0.5))172 173    # 坐标轴174    ax.set(xlabel='Epoch', ylabel='Loss', xlim=(0, len(train_loss) + 1))175 176    # 样式177    ax.legend(frameon=False, loc='upper right')178    ax.grid(True, linestyle='--', alpha=0.3)179    ax.spines[['top', 'right']].set_visible(False)180 181    plt.tight_layout()182    plt.savefig('./result/loss.png', dpi=300, bbox_inches='tight')183    plt.close()184 185 186if __name__ == "__main__":187    current_path = os.getcwd()188    sys.path.append(current_path)189    config_file_path = os.path.join(current_path, 'conf/config.yaml')190    cfg = YParams(config_file_path, 'model')191    cfg_data = YParams(config_file_path, "datapipe")192 193    train_loss = np.load('./data/checkpoints/trloss.npy')194    valid_loss = np.load('./data/checkpoints/valoss.npy')195    plot_loss(train_loss, valid_loss)196 197    data_dir = cfg_data.dataset.data_dir198    total_files, channel_indices, time_step = get_metadata(data_dir, cfg_data.dataset.channels)199 200    # Load data & Compute RMSE/ACC per channel201    h5_files = sorted(glob.glob(os.path.join(data_dir, "data", "*.h5")))202    with h5py.File(h5_files[0], "r") as f:203        mu = f["global_means"][:]204    clim_mean = mu[:, channel_indices, :, :]205    get_result(total_files, channel_indices, time_step, data_dir, clim_mean)206    show_result()207 208    ##### 默认绘制 test_time 第一年的第一个时间步,用户可自行指定日期和变量 #####209    test_year = cfg_data.dataset.test_time[0]210    eg_files = [f'{test_year}010206']211    channel_index = [cfg_data.dataset.channels.index(v) for v in ['2m_temperature', 'geopotential_500', 'temperature_500']]212 213    selected_var = [cfg_data.dataset.channels[int(i)] for i in channel_index]214    print(f"seleted date: {eg_files}")215    print(f"selected channels: {selected_var}")216    for file in eg_files:217        year = file[:4]218        t_idx = filename_to_index(file, time_step)219        with h5py.File(os.path.join(data_dir, 'data', f'{year}.h5'), "r") as f:220            label = f["fields"][t_idx]  # [C, H, W]221            label = label[channel_indices]222        pred = np.load(f'result/output/{file}.npy').squeeze()223        for i in range(len(selected_var)):224            filename = f'./result/{file}_{selected_var[i]}.png'225            plot(label[channel_index[i]], pred[channel_index[i]], selected_var[i], filename)226            print(f'✅plot {filename}')227