Download scripts/result.py from OneScience-Group/XiHe: direct link, hf CLI and curl.
- Browser
- Download file 7.99 kB
-
https://huggingface.co/OneScience-Group/XiHe/resolve/main/scripts/result.py
- Command line
-
hf download hf://OneScience-Group/XiHe/scripts/result.py
-
curl -L -o result.py https://huggingface.co/OneScience-Group/XiHe/resolve/main/scripts/result.py
7.99 kB
| import glob | |
| import os | |
| import sys | |
| from datetime import datetime | |
| import h5py | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| from matplotlib import rcParams | |
| from tqdm import tqdm | |
| from onescience.utils.fcn.YParams import YParams | |
| rcParams["mathtext.fontset"] = "stix" | |
| rcParams["axes.linewidth"] = 0.9 | |
| rcParams["xtick.major.width"] = 0.9 | |
| rcParams["ytick.major.width"] = 0.9 | |
| def get_metadata(data_dir, channels): | |
| 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"]] | |
| time_step = int(ds.attrs["time_step"]) | |
| channel_indices = [all_variables.index(v) for v in channels] | |
| total_files = sorted(f for f in os.listdir("./result/output/") if f.endswith(".npy")) | |
| return total_files, channel_indices, time_step | |
| def filename_to_index(filename, time_step): | |
| dt = datetime.strptime(filename, "%Y%m%d%H") | |
| year_start = datetime(dt.year, 1, 1) | |
| hours = (dt - year_start).total_seconds() / 3600 | |
| return int(hours / time_step) | |
| def get_result(total_files, channel_indices, time_step, data_dir, clim_mean): | |
| channel_rmse = np.zeros(len(channel_indices)) | |
| channel_acc = np.zeros(len(channel_indices)) | |
| clim_mean = clim_mean[0, :, :, :] | |
| if not os.path.exists("./result/rmse.npy") or not os.path.exists("result/acc.npy"): | |
| numerator = np.zeros(len(channel_indices)) | |
| pred_sq_sum = np.zeros(len(channel_indices)) | |
| label_sq_sum = np.zeros(len(channel_indices)) | |
| for file in tqdm(total_files, unit="files"): | |
| fname = file[:-4] | |
| year = fname[:4] | |
| t_idx = filename_to_index(fname, time_step) | |
| with h5py.File(os.path.join(data_dir, "data", f"{year}.h5"), "r") as f: | |
| label = f["fields"][t_idx] | |
| label = label[channel_indices] | |
| pred = np.load(f"result/output/{file}").squeeze() | |
| label_anom = label - clim_mean | |
| pred_anom = pred - clim_mean | |
| numerator += np.sum(pred_anom * label_anom, axis=(1, 2)) | |
| pred_sq_sum += np.sum(pred_anom ** 2, axis=(1, 2)) | |
| label_sq_sum += np.sum(label_anom ** 2, axis=(1, 2)) | |
| channel_rmse += np.sqrt(np.mean((label - pred) ** 2, axis=(1, 2))) | |
| channel_rmse /= len(total_files) | |
| channel_acc = numerator / (np.sqrt(pred_sq_sum * label_sq_sum) + 1e-8) | |
| np.save("./result/acc.npy", channel_acc) | |
| np.save("./result/rmse.npy", channel_rmse) | |
| def show_result(): | |
| channel_rmse = np.load("./result/rmse.npy") | |
| channel_acc = np.load("./result/acc.npy") | |
| channels = [cfg_data.dataset.channels[i] for i in range(len(channel_indices))] | |
| w = 36 | |
| print(f"┌{'─' * (w + 2)}┬{'─' * 14}┬{'─' * 14}┐") | |
| print(f"│ {'Channel':<{w}} │ {'RMSE':>12} │ {'ACC':>12} │") | |
| print(f"├{'─' * (w + 2)}┼{'─' * 14}┼{'─' * 14}┤") | |
| for i, ch in enumerate(channels): | |
| print(f"│ {ch:<{w}} │ {channel_rmse[i]:>12.4f} | {channel_acc[i]:>12.4f} |") | |
| print(f"├{'─' * (w + 2)}┼{'─' * 14}┼{'─' * 14}┤") | |
| print(f"│ {'Average':<{w}} │ {np.mean(channel_rmse):>12.4f} │ {np.mean(channel_acc):>12.4f} │") | |
| print(f"└{'─' * (w + 2)}┴{'─' * 14}┴{'─' * 14}┘") | |
| def plot(label, pred, var, filename): | |
| fig, axes = plt.subplots(1, 3, figsize=(15, 4)) | |
| xtick_labels = ["180°W", "90°W", "0°", "90°E", "180°E"] | |
| ytick_labels = ["90°S", "45°S", "0°", "45°N", "90°N"] | |
| xticks = np.linspace(0, label.shape[-1] - 1, 5) | |
| yticks = np.linspace(0, label.shape[-2] - 1, 5) | |
| vmin = min(label.min(), pred.min()) | |
| vmax = max(label.max(), pred.max()) | |
| diff = label - pred | |
| rmse = np.sqrt(np.mean(diff ** 2)) | |
| diff_abs_max = np.abs(diff).max() | |
| plot_configs = [ | |
| {"data": label, "title": "Truth", "cmap": "viridis", "vmin": vmin, "vmax": vmax}, | |
| {"data": pred, "title": "Prediction", "cmap": "viridis", "vmin": vmin, "vmax": vmax}, | |
| { | |
| "data": diff, | |
| "title": f"Difference (RMSE={rmse:.2f})", | |
| "cmap": "RdBu_r", | |
| "vmin": -diff_abs_max, | |
| "vmax": diff_abs_max, | |
| }, | |
| ] | |
| for ax, cfg in zip(axes, plot_configs): | |
| im = ax.imshow(cfg["data"], cmap=cfg["cmap"], vmin=cfg["vmin"], vmax=cfg["vmax"]) | |
| ax.set_title(cfg["title"], fontsize=12, pad=4) | |
| ax.set_xlabel("Longitude") | |
| ax.set_ylabel("Latitude") | |
| ax.set_xticks(xticks) | |
| ax.set_xticklabels(xtick_labels) | |
| ax.set_yticks(yticks) | |
| ax.set_yticklabels(ytick_labels) | |
| plt.colorbar(im, ax=ax, orientation="horizontal") | |
| fig.suptitle(var, fontsize=14, fontweight="bold", y=0.98) | |
| plt.savefig(filename, dpi=300, bbox_inches="tight") | |
| plt.close() | |
| def plot_loss(train_loss, valid_loss): | |
| mask = ~(np.isnan(train_loss) | np.isnan(valid_loss)) | |
| train_loss = train_loss[mask] | |
| valid_loss = valid_loss[mask] | |
| fig, ax = plt.subplots(figsize=(5, 3.5)) | |
| colors = {"train": "#2563EB", "valid": "#EA580C"} | |
| epochs = np.arange(1, len(train_loss) + 1) | |
| ax.plot(epochs, train_loss, color=colors["train"], linewidth=1.5, label="Train") | |
| ax.plot(epochs, valid_loss, color=colors["valid"], linewidth=1.5, label="Valid", linestyle="--") | |
| min_idx = np.argmin(valid_loss) | |
| ax.scatter(epochs[min_idx], valid_loss[min_idx], color=colors["valid"], s=40, zorder=5, edgecolors="white") | |
| ax.annotate( | |
| f"Best: {valid_loss[min_idx]:.3f}", | |
| xy=(epochs[min_idx], valid_loss[min_idx]), | |
| xytext=(10, 10), | |
| textcoords="offset points", | |
| fontsize=8, | |
| color=colors["valid"], | |
| arrowprops=dict(arrowstyle="-", color=colors["valid"], lw=0.5), | |
| ) | |
| ax.set(xlabel="Epoch", ylabel="Loss", xlim=(0, len(train_loss) + 1)) | |
| ax.legend(frameon=False, loc="upper right") | |
| ax.grid(True, linestyle="--", alpha=0.3) | |
| ax.spines[["top", "right"]].set_visible(False) | |
| plt.tight_layout() | |
| plt.savefig("./result/loss.png", dpi=300, bbox_inches="tight") | |
| plt.close() | |
| if __name__ == "__main__": | |
| current_path = os.getcwd() | |
| sys.path.append(current_path) | |
| config_file_path = os.path.join(current_path, "conf/config.yaml") | |
| cfg = YParams(config_file_path, "model") | |
| cfg_data = YParams(config_file_path, "datapipe") | |
| train_loss = np.load("./data/checkpoints/trloss.npy") | |
| valid_loss = np.load("./data/checkpoints/valoss.npy") | |
| plot_loss(train_loss, valid_loss) | |
| data_dir = cfg_data.dataset.data_dir | |
| total_files, channel_indices, time_step = get_metadata(data_dir, cfg_data.dataset.channels) | |
| mu = np.load(os.path.join(cfg_data.dataset.stats_dir, "global_means.npy")) | |
| clim_mean = mu[:, channel_indices, :, :] | |
| get_result(total_files, channel_indices, time_step, data_dir, clim_mean) | |
| show_result() | |
| test_year = cfg_data.dataset.test_time[0] | |
| eg_files = [f"{test_year}010200"] | |
| channel_index = [ | |
| cfg_data.dataset.channels.index(v) | |
| for v in [ | |
| "sea_surface_height_above_geoid", | |
| "sea_water_potential_temperature_1", | |
| "sea_water_salinity_4", | |
| ] | |
| ] | |
| selected_var = [cfg_data.dataset.channels[int(i)] for i in channel_index] | |
| print(f"seleted date: {eg_files}") | |
| print(f"selected channels: {selected_var}") | |
| for file in eg_files: | |
| year = file[:4] | |
| t_idx = filename_to_index(file, time_step) | |
| with h5py.File(os.path.join(data_dir, "data", f"{year}.h5"), "r") as f: | |
| label = f["fields"][t_idx] | |
| label = label[channel_indices] | |
| pred = np.load(f"result/output/{file}.npy").squeeze() | |
| for i in range(len(selected_var)): | |
| filename = f"./result/{file}_{selected_var[i]}.png" | |
| plot(label[channel_index[i]], pred[channel_index[i]], selected_var[i], filename) | |
| print(f"✅plot {filename}") | |