modified: .DS_Store
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
This commit is contained in:
164
scripts_4/train_tmd_adapter.py
Normal file
164
scripts_4/train_tmd_adapter.py
Normal file
@@ -0,0 +1,164 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user