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:
@@ -1,105 +1,76 @@
|
||||
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
|
||||
import torch
|
||||
|
||||
try:
|
||||
from .config import CORE_FEATURE_NAMES, ExperimentConfig, make_experiment_config
|
||||
from .dataset import build_dataloaders, normalization_from_payload, report_to_text
|
||||
from .model import build_model
|
||||
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, ExperimentConfig, make_experiment_config
|
||||
from dataset import build_dataloaders, normalization_from_payload, report_to_text
|
||||
from model import build_model
|
||||
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 feature MLP on train/val/test splits.")
|
||||
parser.add_argument("--split", choices=("train", "val", "test"), default="val")
|
||||
parser.add_argument("--sample-index", type=int, default=0)
|
||||
parser.add_argument("--all-samples", action="store_true")
|
||||
parser.add_argument("--device", type=str, default=None)
|
||||
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_device(device_name: str | None, config: ExperimentConfig) -> torch.device:
|
||||
requested = device_name or config.train.device
|
||||
if requested.startswith("cuda") and not torch.cuda.is_available():
|
||||
return torch.device("cpu")
|
||||
return torch.device(requested)
|
||||
|
||||
|
||||
def resolve_checkpoint_path(config: ExperimentConfig, checkpoint_arg: str | None) -> Path:
|
||||
def resolve_checkpoint_path(project_root: Path, checkpoint_arg: str | None) -> Path:
|
||||
if checkpoint_arg:
|
||||
return Path(checkpoint_arg).resolve()
|
||||
return config.data.project_root / config.train.checkpoint_dir / "task1_feature_mlp" / config.train.best_model_name
|
||||
return project_root / "checkpoints_final" / "task1_final" / "task1_final_model.pkl"
|
||||
|
||||
|
||||
def relative_percent_error(true_value: float, pred_value: float) -> float:
|
||||
denominator = max(abs(true_value), 1e-12)
|
||||
return abs(pred_value - true_value) / denominator * 100.0
|
||||
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 predict_rms(
|
||||
model: torch.nn.Module,
|
||||
feature_norm: torch.Tensor,
|
||||
x_rms_raw: torch.Tensor,
|
||||
target_mean: torch.Tensor,
|
||||
target_std: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
pred_norm = model(feature_norm)
|
||||
pred_tr = torch.clamp(pred_norm * target_std + target_mean, min=1e-6)
|
||||
return pred_tr * x_rms_raw
|
||||
|
||||
|
||||
def evaluate_sample(
|
||||
model: torch.nn.Module,
|
||||
dataset,
|
||||
raw_record,
|
||||
sample_index: int,
|
||||
device: torch.device,
|
||||
) -> dict[str, float | str]:
|
||||
sample = dataset[sample_index]
|
||||
normalization = dataset.normalization
|
||||
feature_norm = sample["x"].unsqueeze(0).to(device)
|
||||
x_rms_raw = sample["x_rms_raw"].unsqueeze(0).to(device)
|
||||
target_mean = normalization.target_mean.to(device).view(1, 1)
|
||||
target_std = normalization.target_std.to(device).view(1, 1)
|
||||
|
||||
with torch.no_grad():
|
||||
pred_rms = predict_rms(model, feature_norm, x_rms_raw, target_mean, target_std)
|
||||
|
||||
true_rms = float(sample["y_rms_raw"].item())
|
||||
pred_rms_value = float(pred_rms.detach().cpu().numpy().reshape(-1)[0])
|
||||
row = {
|
||||
"file_name": str(sample["file_name"]),
|
||||
"frequency_hz": float(sample["frequency_hz"].item()),
|
||||
"true_rms": true_rms,
|
||||
"pred_rms": pred_rms_value,
|
||||
"relative_error_percent": relative_percent_error(true_rms, pred_rms_value),
|
||||
}
|
||||
features_raw = sample["features_raw"].numpy().reshape(-1)
|
||||
for name, value in zip(CORE_FEATURE_NAMES, features_raw):
|
||||
row[name] = float(value)
|
||||
row["sampling_rate"] = float(raw_record.sampling_rate)
|
||||
return row
|
||||
|
||||
|
||||
def save_all_samples_plot(result_df: pd.DataFrame, split: str, save_dir: Path) -> Path | None:
|
||||
if result_df["frequency_hz"].isna().any():
|
||||
return None
|
||||
def build_dense_curve(model, result_df: pd.DataFrame, config) -> pd.DataFrame:
|
||||
plot_df = result_df.sort_values("frequency_hz").reset_index(drop=True)
|
||||
fig, axes = plt.subplots(2, 1, figsize=(12, 8))
|
||||
fig.suptitle(f"Task1 Feature-MLP Evaluation | {split}")
|
||||
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,
|
||||
}
|
||||
)
|
||||
|
||||
axes[0].plot(plot_df["frequency_hz"], plot_df["true_rms"], marker="o", label="True RMS")
|
||||
axes[0].plot(plot_df["frequency_hz"], plot_df["pred_rms"], marker="o", label="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)
|
||||
@@ -112,112 +83,58 @@ def save_all_samples_plot(result_df: pd.DataFrame, split: str, save_dir: Path) -
|
||||
axes[1].tick_params(axis="x", labelrotation=45)
|
||||
|
||||
plt.tight_layout()
|
||||
save_dir.mkdir(parents=True, exist_ok=True)
|
||||
figure_path = save_dir / f"evaluation_{split}_curve.png"
|
||||
plt.savefig(figure_path, dpi=180, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
return figure_path
|
||||
|
||||
|
||||
def save_single_sample_plot(raw_record, result: dict[str, float | str], split: str, save_dir: Path, sample_index: int) -> Path:
|
||||
time_middle = raw_record.time_middle.detach().cpu().numpy().reshape(-1)
|
||||
x_middle = raw_record.x_middle.detach().cpu().numpy().reshape(-1)
|
||||
y_middle = raw_record.y_middle.detach().cpu().numpy().reshape(-1)
|
||||
frequency_hz = float(result["frequency_hz"])
|
||||
pred_rms = float(result["pred_rms"])
|
||||
true_rms = float(result["true_rms"])
|
||||
|
||||
fig, axes = plt.subplots(3, 1, figsize=(12, 10))
|
||||
fig.suptitle(f"Task1 Feature-MLP | {split} | {result['file_name']}")
|
||||
|
||||
axes[0].plot(time_middle, x_middle, color="tab:blue")
|
||||
axes[0].set_title("Input Base Excitation (Middle Segment)")
|
||||
axes[0].set_xlabel("Time")
|
||||
axes[0].set_ylabel("Acceleration")
|
||||
axes[0].grid(True, alpha=0.3)
|
||||
|
||||
axes[1].plot(time_middle, y_middle, color="tab:green")
|
||||
axes[1].set_title("True Top Response (Middle Segment)")
|
||||
axes[1].set_xlabel("Time")
|
||||
axes[1].set_ylabel("Acceleration")
|
||||
axes[1].grid(True, alpha=0.3)
|
||||
|
||||
axes[2].bar(["True RMS", "Pred RMS"], [true_rms, pred_rms], color=["tab:green", "tab:orange"])
|
||||
axes[2].set_title(
|
||||
f"Freq: {frequency_hz:.4f} Hz | True RMS: {true_rms:.6f} | Pred RMS: {pred_rms:.6f} | "
|
||||
f"Error: {float(result['relative_error_percent']):.2f}%"
|
||||
)
|
||||
axes[2].set_ylabel("RMS")
|
||||
axes[2].grid(True, axis="y", alpha=0.3)
|
||||
|
||||
plt.tight_layout()
|
||||
save_dir.mkdir(parents=True, exist_ok=True)
|
||||
figure_path = save_dir / f"evaluation_{split}_s{sample_index}.png"
|
||||
plt.savefig(figure_path, dpi=180, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
return figure_path
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
config = make_experiment_config()
|
||||
device = resolve_device(args.device, config)
|
||||
checkpoint_path = resolve_checkpoint_path(config, args.checkpoint)
|
||||
if not checkpoint_path.exists():
|
||||
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
|
||||
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"]
|
||||
|
||||
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
if "normalization" in checkpoint:
|
||||
normalization = normalization_from_payload(checkpoint["normalization"])
|
||||
else:
|
||||
normalization = None
|
||||
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()
|
||||
|
||||
loaders, datasets, raw_records, reports = build_dataloaders(config)
|
||||
del loaders
|
||||
print(report_to_text(reports))
|
||||
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)
|
||||
|
||||
dataset = datasets[args.split]
|
||||
if normalization is not None:
|
||||
dataset.normalization = normalization
|
||||
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)
|
||||
|
||||
model = build_model(config).to(device)
|
||||
model.load_state_dict(checkpoint["model_state_dict"])
|
||||
model.eval()
|
||||
|
||||
save_dir = config.data.project_root / "evaluation_outputs" / "task1_feature_mlp"
|
||||
if args.all_samples:
|
||||
rows = [
|
||||
evaluate_sample(model, dataset, raw_records[args.split][index], index, device)
|
||||
for index in range(len(dataset))
|
||||
]
|
||||
result_df = pd.DataFrame(rows).sort_values(["relative_error_percent", "file_name"]).reset_index(drop=True)
|
||||
save_dir.mkdir(parents=True, exist_ok=True)
|
||||
csv_path = save_dir / f"evaluation_{args.split}_all_samples.csv"
|
||||
result_df.to_csv(csv_path, index=False)
|
||||
figure_path = save_all_samples_plot(result_df, args.split, save_dir)
|
||||
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: {checkpoint_path}")
|
||||
print(f"Summary: {summary}")
|
||||
print(f"CSV saved to: {csv_path}")
|
||||
if figure_path is not None:
|
||||
print(f"Figure saved to: {figure_path}")
|
||||
print(result_df.to_string(index=False))
|
||||
return
|
||||
|
||||
result = evaluate_sample(model, dataset, raw_records[args.split][args.sample_index], args.sample_index, device)
|
||||
figure_path = save_single_sample_plot(raw_records[args.split][args.sample_index], result, args.split, save_dir, args.sample_index)
|
||||
print(f"Checkpoint: {checkpoint_path}")
|
||||
print(f"Sample file: {result['file_name']}")
|
||||
print(f"True RMS: {float(result['true_rms']):.6f}")
|
||||
print(f"Pred RMS: {float(result['pred_rms']):.6f}")
|
||||
print(f"Relative RMS Error (%): {float(result['relative_error_percent']):.4f}")
|
||||
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__":
|
||||
|
||||
Reference in New Issue
Block a user