from __future__ import annotations from functools import lru_cache from pathlib import Path import numpy as np import pandas as pd def _task1_curve_candidates(project_root: Path) -> tuple[Path, ...]: base_dir = project_root / "evaluation_outputs" / "task1_final" return ( base_dir / "task1_final_dense_curve.csv", base_dir / "evaluation_dense_curve.csv", ) @lru_cache(maxsize=8) def load_task1_dense_curve(project_root_str: str) -> tuple[np.ndarray, np.ndarray, np.ndarray] | None: project_root = Path(project_root_str) for path in _task1_curve_candidates(project_root): if not path.exists(): continue try: curve_df = pd.read_csv(path) except Exception: continue required = {"frequency_hz", "pred_tr", "x_rms_interp"} if not required.issubset(set(curve_df.columns)): continue cleaned = curve_df[["frequency_hz", "pred_tr", "x_rms_interp"]].copy() cleaned["frequency_hz"] = pd.to_numeric(cleaned["frequency_hz"], errors="coerce") cleaned["pred_tr"] = pd.to_numeric(cleaned["pred_tr"], errors="coerce") cleaned["x_rms_interp"] = pd.to_numeric(cleaned["x_rms_interp"], errors="coerce") cleaned = cleaned.dropna() if cleaned.empty: continue cleaned = cleaned.sort_values("frequency_hz").drop_duplicates(subset="frequency_hz", keep="first") return ( cleaned["frequency_hz"].to_numpy(dtype=np.float64), cleaned["pred_tr"].to_numpy(dtype=np.float64), cleaned["x_rms_interp"].to_numpy(dtype=np.float64), ) return None def predict_task1_tr( model, *, feature_vector: np.ndarray, frequency_hz: float, normalization_eps: float, project_root: Path, prefer_curve: bool = True, ) -> tuple[float, str]: features = np.asarray(feature_vector, dtype=np.float64).reshape(1, -1) tr_value = max(float(model.predict(features)[0]), normalization_eps) source = "model_direct" if prefer_curve: curve = load_task1_dense_curve(str(project_root.resolve())) if curve is not None and features.shape[1] >= 3: freq_grid, _tr_grid, x_ref_grid = curve query_freq = float(np.clip(frequency_hz, freq_grid.min(), freq_grid.max())) x_ref = float(np.interp(query_freq, freq_grid, x_ref_grid)) x_obs = float(features[0, 2]) ratio = x_obs / max(x_ref, normalization_eps) if ratio > 1.15 or ratio < 0.85: tr_value = max(tr_value / ratio, normalization_eps) source = "model_amp_adapt" return tr_value, source