modified: checkpoints_final/task1_final/task1_final_model.pkl new file: checkpoints_final/task1_final/task1_tmd_adapter.pkl new file: scripts_4/README.md new file: scripts_4/__init__.py new file: scripts_4/__pycache__/adapter.cpython-314.pyc new file: scripts_4/adapter.py new file: scripts_4/predict_tmd_single.py new file: scripts_4/train_tmd_adapter.py
165 lines
6.7 KiB
Python
165 lines
6.7 KiB
Python
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import pickle
|
|
from pathlib import Path
|
|
import sys
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
|
|
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
|
sys.path.insert(0, str(PROJECT_ROOT / "scripts"))
|
|
|
|
from config import make_experiment_config
|
|
from dataset import build_record_from_file
|
|
|
|
try:
|
|
from .adapter import FrequencyScaleAdapter, save_adapter
|
|
except ImportError:
|
|
from adapter import FrequencyScaleAdapter, save_adapter
|
|
|
|
|
|
def parse_args() -> argparse.Namespace:
|
|
parser = argparse.ArgumentParser(description="Fit a tiny-sample TMD adapter on top of the Non_TMD baseline .pkl model.")
|
|
parser.add_argument("--checkpoint", type=str, default=None, help="Path to baseline model .pkl.")
|
|
parser.add_argument("--tmd-dir", type=str, default="downloads/TMD", help="Directory that contains TMD harmonic csv files.")
|
|
parser.add_argument("--output", type=str, default="checkpoints_final/task1_final/task1_tmd_adapter.pkl")
|
|
parser.add_argument("--report-json", type=str, default="evaluation_outputs/task1_final/tmd_adapter_report.json")
|
|
parser.add_argument("--shrinkage", type=float, default=0.20, help="Scale shrinkage toward 1.0.")
|
|
parser.add_argument("--extrapolation-decay-hz", type=float, default=0.25)
|
|
parser.add_argument("--min-scale", type=float, default=0.40)
|
|
parser.add_argument("--max-scale", type=float, default=2.50)
|
|
return parser.parse_args()
|
|
|
|
|
|
def resolve_checkpoint_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 load_baseline_model(model_path: Path):
|
|
with model_path.open("rb") as handle:
|
|
payload = pickle.load(handle)
|
|
return payload["model"], payload.get("summary", {})
|
|
|
|
|
|
def base_predict_rms(model, features: np.ndarray, x_rms: float, eps: float) -> float:
|
|
pred_tr = max(float(model.predict([features.tolist()])[0]), eps)
|
|
return pred_tr * float(x_rms)
|
|
|
|
|
|
def collect_tmd_rows(model, tmd_dir: Path, eps: float) -> list[dict]:
|
|
config = make_experiment_config().data
|
|
config.data_dir = str(Path(tmd_dir).relative_to(PROJECT_ROOT))
|
|
config.__post_init__()
|
|
rows: list[dict] = []
|
|
for file_path in sorted(config.data_root.rglob("harmonic_5mm_*Hz_TMD.csv")):
|
|
record = build_record_from_file(file_path, config)
|
|
base_pred = base_predict_rms(model, record.features, record.x_rms, eps)
|
|
scale = float(record.y_rms / max(base_pred, eps))
|
|
rows.append(
|
|
{
|
|
"file_name": file_path.name,
|
|
"frequency_hz": float(record.frequency_hz),
|
|
"true_y_rms": float(record.y_rms),
|
|
"base_pred_y_rms": float(base_pred),
|
|
"scale_true_over_base": float(scale),
|
|
}
|
|
)
|
|
if not rows:
|
|
raise RuntimeError("No TMD harmonic files found.")
|
|
return rows
|
|
|
|
|
|
def loocv_error(rows: list[dict], shrinkage: float, extrapolation_decay_hz: float, min_scale: float, max_scale: float) -> list[float]:
|
|
errors: list[float] = []
|
|
eps = 1e-12
|
|
for holdout in range(len(rows)):
|
|
train_rows = [row for i, row in enumerate(rows) if i != holdout]
|
|
adapter = FrequencyScaleAdapter(
|
|
frequencies_hz=np.asarray([row["frequency_hz"] for row in train_rows], dtype=np.float64),
|
|
scale_values=np.asarray([row["scale_true_over_base"] for row in train_rows], dtype=np.float64),
|
|
shrinkage=shrinkage,
|
|
extrapolation_decay_hz=extrapolation_decay_hz,
|
|
min_scale=min_scale,
|
|
max_scale=max_scale,
|
|
)
|
|
sample = rows[holdout]
|
|
pred = adapter.predict(sample["base_pred_y_rms"], sample["frequency_hz"])
|
|
err = abs(pred - sample["true_y_rms"]) / max(abs(sample["true_y_rms"]), eps) * 100.0
|
|
errors.append(float(err))
|
|
return errors
|
|
|
|
|
|
def main() -> None:
|
|
args = parse_args()
|
|
ckpt_path = resolve_checkpoint_path(args.checkpoint)
|
|
model, baseline_summary = load_baseline_model(ckpt_path)
|
|
rows = collect_tmd_rows(model, (PROJECT_ROOT / args.tmd_dir).resolve(), eps=1e-6)
|
|
rows = sorted(rows, key=lambda x: x["frequency_hz"])
|
|
|
|
adapter = FrequencyScaleAdapter(
|
|
frequencies_hz=np.asarray([row["frequency_hz"] for row in rows], dtype=np.float64),
|
|
scale_values=np.asarray([row["scale_true_over_base"] for row in rows], dtype=np.float64),
|
|
shrinkage=args.shrinkage,
|
|
extrapolation_decay_hz=args.extrapolation_decay_hz,
|
|
min_scale=args.min_scale,
|
|
max_scale=args.max_scale,
|
|
)
|
|
|
|
adapted_rows = []
|
|
for row in rows:
|
|
adapted_pred = adapter.predict(row["base_pred_y_rms"], row["frequency_hz"])
|
|
adapted_err = abs(adapted_pred - row["true_y_rms"]) / max(abs(row["true_y_rms"]), 1e-12) * 100.0
|
|
adapted_rows.append({**row, "adapted_pred_y_rms": adapted_pred, "adapted_err_pct": adapted_err})
|
|
|
|
loocv_errors = loocv_error(
|
|
rows=rows,
|
|
shrinkage=args.shrinkage,
|
|
extrapolation_decay_hz=args.extrapolation_decay_hz,
|
|
min_scale=args.min_scale,
|
|
max_scale=args.max_scale,
|
|
)
|
|
report = {
|
|
"checkpoint": str(ckpt_path),
|
|
"baseline_summary": baseline_summary,
|
|
"samples": adapted_rows,
|
|
"adapter_scales_after_shrinkage": [
|
|
{"frequency_hz": float(f), "scale": float(s)}
|
|
for f, s in zip(adapter.frequencies_hz.tolist(), adapter.scale_values.tolist())
|
|
],
|
|
"metrics": {
|
|
"base_mean_err_pct": float(
|
|
np.mean(
|
|
[
|
|
abs(row["base_pred_y_rms"] - row["true_y_rms"]) / max(abs(row["true_y_rms"]), 1e-12) * 100.0
|
|
for row in rows
|
|
]
|
|
)
|
|
),
|
|
"adapted_mean_err_pct": float(np.mean([row["adapted_err_pct"] for row in adapted_rows])),
|
|
"loocv_mean_err_pct": float(np.mean(loocv_errors)),
|
|
"loocv_max_err_pct": float(np.max(loocv_errors)),
|
|
},
|
|
}
|
|
|
|
output_path = (PROJECT_ROOT / args.output).resolve()
|
|
save_adapter(adapter, output_path, extra={"report_metrics": report["metrics"], "checkpoint": str(ckpt_path)})
|
|
|
|
report_path = (PROJECT_ROOT / args.report_json).resolve()
|
|
report_path.parent.mkdir(parents=True, exist_ok=True)
|
|
report_path.write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8")
|
|
|
|
print(f"Adapter saved to: {output_path}")
|
|
print(f"Report saved to: {report_path}")
|
|
print(pd.DataFrame(adapted_rows).to_string(index=False))
|
|
print(f"LOOCV mean error (%): {report['metrics']['loocv_mean_err_pct']:.4f}")
|
|
print(f"LOOCV max error (%): {report['metrics']['loocv_max_err_pct']:.4f}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|