modified: .gitignore deleted: Figure_1.png deleted: checkpoints/forward/best_tcn_model.pt deleted: checkpoints/forward/training_history.csv deleted: checkpoints_mlp/task1_feature_mlp/best_feature_mlp.pt deleted: checkpoints_mlp/task1_feature_mlp/training_history.csv deleted: checkpoints_rms/forward_rms/best_rms_model.pt deleted: checkpoints_rms/forward_rms/training_history.csv deleted: checkpoints_tree/task1_tr_tree/best_tr_tree.pkl deleted: checkpoints_tree/task1_tr_tree/model_selection.csv deleted: evaluation_outputs/forward/evaluation_forward_test_b0_s0.png deleted: evaluation_outputs/forward/evaluation_forward_val_b0_s0.png deleted: evaluation_outputs/forward/evaluation_forward_val_b3_s0.png deleted: evaluation_outputs/forward_rms/evaluation_test_all_samples.csv deleted: evaluation_outputs/forward_rms/evaluation_test_s0.png deleted: evaluation_outputs/forward_rms/evaluation_train_all_samples.csv deleted: evaluation_outputs/forward_rms/evaluation_val_all_samples.csv deleted: evaluation_outputs/forward_rms/evaluation_val_s0.png deleted: evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png deleted: evaluation_outputs/task1_feature_mlp/evaluation_train_all_samples.csv deleted: evaluation_outputs/task1_feature_mlp/evaluation_train_curve.png deleted: evaluation_outputs/task1_feature_mlp/evaluation_val_all_samples.csv deleted: evaluation_outputs/task1_feature_mlp/evaluation_val_curve.png deleted: evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png deleted: evaluation_outputs/task1_feature_mlp/harmonic_5mm_0.75Hz_prediction.png deleted: evaluation_outputs/task1_feature_mlp/harmonic_5mm_1.55Hz_prediction.png deleted: evaluation_outputs/task1_tr_tree/evaluation_train_all_samples.csv deleted: evaluation_outputs/task1_tr_tree/evaluation_train_curve.png deleted: evaluation_outputs/task1_tr_tree/evaluation_val_all_samples.csv deleted: evaluation_outputs/task1_tr_tree/evaluation_val_curve.png deleted: evaluation_outputs/task1_tr_tree/harmonic_5mm_1.55Hz_prediction.png deleted: sanity_check_alignment_forward.png new file: scripts/README.md modified: scripts/__pycache__/config.cpython-310.pyc modified: scripts/__pycache__/dataset.cpython-310.pyc deleted: scripts/__pycache__/model.cpython-310.pyc modified: scripts/config.py modified: scripts/dataset.py modified: scripts/evaluate.py deleted: scripts/model.py modified: scripts/predict_single.py deleted: scripts/train.py new file: scripts/train_final.py deleted: scripts_tree/__pycache__/config.cpython-310.pyc deleted: scripts_tree/config.py deleted: scripts_tree/evaluate.py deleted: scripts_tree/predict_single.py deleted: scripts_tree/train.py deleted: src/__pycache__/config.cpython-310.pyc deleted: src/__pycache__/config.cpython-314.pyc deleted: src/__pycache__/dataset.cpython-310.pyc deleted: src/__pycache__/dataset.cpython-314.pyc deleted: src/__pycache__/model.cpython-310.pyc deleted: src/__pycache__/model.cpython-314.pyc deleted: src/config.py deleted: src/dataset.py deleted: src/evaluate.py deleted: src/model.py deleted: src/sanity_check.py deleted: src/train.py deleted: src_new/__pycache__/config.cpython-310.pyc deleted: src_new/__pycache__/dataset.cpython-310.pyc deleted: src_new/__pycache__/evaluate.cpython-310.pyc deleted: src_new/__pycache__/model.cpython-310.pyc deleted: src_new/__pycache__/train.cpython-310.pyc deleted: src_new/config.py deleted: src_new/dataset.py deleted: src_new/evaluate.py deleted: src_new/model.py deleted: src_new/train.py deleted: src_old/__init__.py deleted: src_old/__pycache__/config.cpython-310.pyc deleted: src_old/__pycache__/config.cpython-314.pyc deleted: src_old/__pycache__/dataset.cpython-310.pyc deleted: src_old/__pycache__/dataset.cpython-314.pyc deleted: src_old/__pycache__/evaluate.cpython-314.pyc deleted: src_old/__pycache__/model.cpython-310.pyc deleted: src_old/__pycache__/model.cpython-314.pyc deleted: src_old/__pycache__/train.cpython-310.pyc deleted: src_old/__pycache__/train.cpython-314.pyc deleted: src_old/config.py deleted: src_old/dataset.py deleted: src_old/evaluate.py deleted: src_old/model.py deleted: src_old/train.py
179 lines
6.7 KiB
Python
179 lines
6.7 KiB
Python
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()
|