Files
Building/scripts/evaluate.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

142 lines
5.7 KiB
Python

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