from __future__ import annotations import argparse import pickle from pathlib import Path import matplotlib.pyplot as plt import numpy as np import pandas as pd try: from .config import CORE_FEATURE_NAMES, evaluation_dir, make_experiment_config from .dataset import load_all_records, records_to_frame, report_to_text except ImportError: from config import CORE_FEATURE_NAMES, evaluation_dir, make_experiment_config from dataset import load_all_records, records_to_frame, report_to_text def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Evaluate the final all-data task1 model.") parser.add_argument("--checkpoint", type=str, default=None) return parser.parse_args() def resolve_checkpoint_path(project_root: Path, checkpoint_arg: str | None) -> Path: if checkpoint_arg: return Path(checkpoint_arg).resolve() return project_root / "checkpoints_final" / "task1_final" / "task1_final_model.pkl" def relative_error_percent(true_value: float, pred_value: float) -> float: return abs(pred_value - true_value) / max(abs(true_value), 1e-12) * 100.0 def build_dense_curve(model, result_df: pd.DataFrame, config) -> pd.DataFrame: plot_df = result_df.sort_values("frequency_hz").reset_index(drop=True) dense_freq = np.linspace( float(plot_df["frequency_hz"].min()), float(plot_df["frequency_hz"].max()), config.train.dense_curve_points, ) dense_x_rms = np.interp(dense_freq, plot_df["frequency_hz"], plot_df["x_rms"]) dense_true_rms = np.interp(dense_freq, plot_df["frequency_hz"], plot_df["true_rms"]) dense_features = np.column_stack( [ dense_freq, dense_freq**2, dense_x_rms, (2.0 * np.pi * dense_freq) ** 2 * config.data.harmonic_amplitude_m, ] ) dense_pred_tr = model.predict(dense_features).clip(min=config.data.normalization_eps) dense_pred_rms = dense_pred_tr * dense_x_rms return pd.DataFrame( { "frequency_hz": dense_freq, "x_rms_interp": dense_x_rms, "true_rms_interp": dense_true_rms, "pred_tr": dense_pred_tr, "pred_rms": dense_pred_rms, } ) def save_plot(result_df: pd.DataFrame, dense_df: pd.DataFrame, figure_path: Path) -> None: plot_df = result_df.sort_values("frequency_hz").reset_index(drop=True) fig, axes = plt.subplots(2, 1, figsize=(12, 9)) fig.suptitle("Task1 Final Model Evaluation | All Harmonic Data") axes[0].plot(dense_df["frequency_hz"], dense_df["true_rms_interp"], color="tab:blue", linewidth=2.0, label="True RMS Guide Curve") axes[0].plot(dense_df["frequency_hz"], dense_df["pred_rms"], color="tab:orange", linewidth=2.0, label="Pred RMS Dense Curve") axes[0].scatter(plot_df["frequency_hz"], plot_df["true_rms"], color="tab:blue", s=28, zorder=3, label="True RMS Samples") axes[0].scatter(plot_df["frequency_hz"], plot_df["pred_rms"], color="tab:orange", s=22, zorder=3, label="Pred RMS Samples") axes[0].set_xlabel("Frequency (Hz)") axes[0].set_ylabel("RMS") axes[0].grid(True, alpha=0.3) axes[0].legend() axes[1].bar(plot_df["frequency_hz"].astype(str), plot_df["relative_error_percent"], color="tab:orange") axes[1].set_xlabel("Frequency (Hz)") axes[1].set_ylabel("Relative Error (%)") axes[1].grid(True, axis="y", alpha=0.3) axes[1].tick_params(axis="x", labelrotation=45) plt.tight_layout() plt.savefig(figure_path, dpi=180, bbox_inches="tight") plt.close(fig) def main() -> None: args = parse_args() config = make_experiment_config() ckpt_path = resolve_checkpoint_path(config.data.project_root, args.checkpoint) with ckpt_path.open("rb") as handle: payload = pickle.load(handle) model = payload["model"] records, report = load_all_records(config) print(report_to_text(report)) data_df = records_to_frame(records) feature_matrix = data_df.loc[:, CORE_FEATURE_NAMES].to_numpy() pred_tr = model.predict(feature_matrix).clip(min=config.data.normalization_eps) pred_rms = pred_tr * data_df["x_rms"].to_numpy() true_rms = data_df["y_rms"].to_numpy() result_df = data_df.copy() result_df["pred_tr"] = pred_tr result_df["pred_rms"] = pred_rms result_df["true_rms"] = true_rms result_df["relative_error_percent"] = [ relative_error_percent(true, pred) for true, pred in zip(true_rms, pred_rms) ] result_df = result_df.sort_values(["frequency_hz", "file_name"]).reset_index(drop=True) save_dir = evaluation_dir(config) csv_path = save_dir / "evaluation_all_samples.csv" dense_csv_path = save_dir / config.train.evaluation_dense_csv_name figure_path = save_dir / "evaluation_all_curve.png" dense_df = build_dense_curve(model, result_df, config) result_df.to_csv(csv_path, index=False) dense_df.to_csv(dense_csv_path, index=False) save_plot(result_df, dense_df, figure_path) summary = { "count": len(result_df), "mean_error_percent": float(result_df["relative_error_percent"].mean()), "median_error_percent": float(result_df["relative_error_percent"].median()), "max_error_percent": float(result_df["relative_error_percent"].max()), "min_error_percent": float(result_df["relative_error_percent"].min()), } print(f"Checkpoint: {ckpt_path}") print(f"Model: {payload['model_name']}") print(f"Summary: {summary}") print(f"CSV saved to: {csv_path}") print(f"Dense curve CSV saved to: {dense_csv_path}") print(f"Figure saved to: {figure_path}") print(result_df[["file_name", "frequency_hz", "true_rms", "pred_rms", "relative_error_percent"]].to_string(index=False)) if __name__ == "__main__": main()