from __future__ import annotations import pickle from dataclasses import asdict from pathlib import Path from typing import Any import matplotlib.pyplot as plt import numpy as np import pandas as pd from sklearn.neighbors import KNeighborsRegressor from sklearn.pipeline import Pipeline from sklearn.preprocessing import StandardScaler try: from .config import CORE_FEATURE_NAMES, ExperimentConfig, checkpoint_dir, 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, ExperimentConfig, checkpoint_dir, evaluation_dir, make_experiment_config from dataset import load_all_records, records_to_frame, report_to_text def build_final_model(config: ExperimentConfig) -> Pipeline: return Pipeline( [ ("scaler", StandardScaler()), ( "model", KNeighborsRegressor( n_neighbors=config.model.n_neighbors, weights="distance", p=config.model.distance_power, ), ), ] ) def serialize_for_checkpoint(value: Any) -> Any: if isinstance(value, Path): return str(value) if isinstance(value, dict): return {key: serialize_for_checkpoint(sub_value) for key, sub_value in value.items()} if isinstance(value, tuple): return [serialize_for_checkpoint(item) for item in value] if isinstance(value, list): return [serialize_for_checkpoint(item) for item in value] return value 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: Pipeline, result_df: pd.DataFrame, config: ExperimentConfig) -> 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_fit_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 | All Harmonic Data Fit") 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 train_final() -> None: config = make_experiment_config() 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() target_tr = data_df["target_tr"].to_numpy() model = build_final_model(config) model.fit(feature_matrix, target_tr) 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) mean_error = float(result_df["relative_error_percent"].mean()) median_error = float(result_df["relative_error_percent"].median()) max_error = float(result_df["relative_error_percent"].max()) ckpt_dir = checkpoint_dir(config) eval_dir = evaluation_dir(config) csv_path = eval_dir / config.train.fit_csv_name fig_path = eval_dir / config.train.fit_figure_name dense_csv_path = eval_dir / config.train.dense_curve_csv_name model_path = ckpt_dir / config.train.model_name 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_fit_plot(result_df, dense_df, fig_path) payload = { "model_name": config.model.model_name, "model": model, "config": serialize_for_checkpoint(asdict(config)), "feature_names": list(CORE_FEATURE_NAMES), "summary": { "count": len(result_df), "mean_error_percent": mean_error, "median_error_percent": median_error, "max_error_percent": max_error, }, } with model_path.open("wb") as handle: pickle.dump(payload, handle) print(f"Final model saved to: {model_path}") print(f"Fit CSV saved to: {csv_path}") print(f"Dense curve CSV saved to: {dense_csv_path}") print(f"Fit figure saved to: {fig_path}") print( f"All-data self-fit summary | count={len(result_df)} | " f"mean={mean_error:.4f}% | median={median_error:.4f}% | max={max_error:.4f}%" ) print(result_df[["file_name", "frequency_hz", "true_rms", "pred_rms", "relative_error_percent"]].to_string(index=False)) if __name__ == "__main__": train_final()