XiHe / scripts /result.py
OneScience's picture
Upload folder using huggingface_hub
2429091 verified
Raw
History Blame Contribute Delete
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}")