OneScience-Group/Pangu_Weather
038
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 