Files
Building/scripts/train_final.py
CrbnsCat10n fc8fcc2746 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
2026-05-06 14:37:37 +08:00

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()