Changes to be committed:
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
This commit is contained in:
178
scripts/train_final.py
Normal file
178
scripts/train_final.py
Normal file
@@ -0,0 +1,178 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user