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