From 79f3de05f7cd8f8326b16d98a37fcf2b45dbc334 Mon Sep 17 00:00:00 2001 From: CrbnsCat10n Date: Wed, 6 May 2026 20:57:43 +0800 Subject: [PATCH] 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 --- .DS_Store | Bin 8196 -> 8196 bytes .../task1_final/task1_final_model.pkl | Bin 4787 -> 4783 bytes .../task1_final/task1_tmd_adapter.pkl | Bin 0 -> 583 bytes scripts_4/README.md | 45 +++++ scripts_4/__init__.py | 2 + scripts_4/__pycache__/adapter.cpython-314.pyc | Bin 0 -> 8569 bytes scripts_4/adapter.py | 101 +++++++++++ scripts_4/predict_tmd_single.py | 78 +++++++++ scripts_4/train_tmd_adapter.py | 164 ++++++++++++++++++ 9 files changed, 390 insertions(+) create mode 100644 checkpoints_final/task1_final/task1_tmd_adapter.pkl create mode 100644 scripts_4/README.md create mode 100644 scripts_4/__init__.py create mode 100644 scripts_4/__pycache__/adapter.cpython-314.pyc create mode 100644 scripts_4/adapter.py create mode 100644 scripts_4/predict_tmd_single.py create mode 100644 scripts_4/train_tmd_adapter.py diff --git a/.DS_Store b/.DS_Store index e677bc4c046be0c294a283d7a3d902ec61818a9c..c9da38d6aeecd29fce375b62503363ff048fbf3e 100644 GIT binary patch delta 56 zcmZp1XmOa}&&a(oU^hP__hudeeHLL(hGK?fh9ZVch608XAUmGHWU`Kk?dB;m10!-HP0Wj* z(jx>>mYiRdst07I7G)++>EVt~E6UGJDosmEEt)cUipDerux18^$y*r}6*J_#Sz4#` zFgi~u0hz>7mRJ-&B|{xa)i1cJ3@xx4w&LW(oK!Sbj!cSb5VN4F1mcUcb5awF^5V-< zi;6Sz^MKAY)U(htm|V|vl@n-z2P4QGoAa2x8F_#@HCv}ZbWYyG5)YDBUyoRiClc!{~Pfp-f z^8;$(farjeQ-E4f%*3j|n-Lal?NgHcfbPgJ&#;($k5`WwMDz1$=<4W)76YS5Ke;F= z&qP0{gnLjyqx^Rlwy6q{Ji)O YUzaJR8O*i77-P&}%CO(OMu43O0Ix;!f&c&j delta 765 zcmZ3lx>=Q_fo1B7i7d}J#PUjW3o7;ECw|mt%$%&qs9dkf00C3lrUXstVa-iV%!{9r zA=bkaUzD1hpI2N`RGM5eW%85^$sYE+l*FQ<#7ZE~o1t}zGh@<}cAzE=Z$@vH)+rg% zJxpm9Q#v~!V$2!xetv#l|A7EZcr%ntN$PY~fY|oKK5)0t*)wNYCQoM+QUS{Ju#_bh z0WDU`&;Xg8p@nR+0Z`xMQ;dpgAcbtj$%#3sc$M-nDawMB3d9#@=cFbU<;9n!78Pga z=K(!vsAsNcG&z{*suYG*nm~IEI-NWiHybj0Gcp5BnB2hk zuu4TP7rmSwHgA`ZqSRDi2(^MjCL;~gLm*q^z_v7Tc{7RwZNXH-*g9FA$6co}L(ZE4 zXjK!C<^a;oKw1Duw`3T2OSDZX^;Q4|y#UYzQ24O(C@Nw#j0qG!?H~s`J4`;$;|O-H z0Gng{GHm?w1X956fWJ(h7Y`^8LmGP|80!L$h8+R{0W z2ffv!Cr|zz9yJ#W6%4c#DLsgX(n%((kbb8(@4fHMH{a|Q_tqX1#QCSizxiL#P#QDGYYTIs=X)Qk6q#rSJ2FD6#ja4IJhV+f|-*Y zL3x6=z(zi^P%w{STM3+{3@e&)6a&Hi(RfjK>OE$zOSylWBhU9ob>K+6os*53iL;EE z035qbbcE)K{iIEC9Pbw(AWMZpT~LA&Ip!)YEWt6DBM<^*MA@jZTK@g;<;rU_@7EVk z>!r(+sOs33?5u1qC5|N86v@mC z6U*3L2t|P^#TId07p8*-R@(q816f2VwvG1Ju7O3@1^R>Bk(SI<3mDiGi=b$sz%dM@ zKiYHdaE2q5vbWtU^Ui&od*{xX^PSiDfTzkyL0M?{ulSF=6!jbI=!Mf7^wR}sEKt`d zo*t&o(lXtOZDyD`YmqIaZ5d|IaxzES?66g~!aO@{lkGg$K=IZFD&dnIS1ol^8^zmT z#NNwjd-rM=N_O&9yaQ&ONSmN#7w>|eo2+#6YzGzge2KJST6IN}$&?(G0#h5 z?3Z{lm3t?e$OtMook&Gx)hY^dMod1$yIZGI;#^dgaQk+3sOOR-h*Hm#IGOD285L94 zg(+F;IhlzkV)5imWDKVawrATj2G!{+Eg4ms&SKX0F>4AhuMD)P6_(r4_ zs1~T_6ivTIe?{usnN4zCs?+*Ufr28BiVZ9__>o8b0W=n<37rMhM2X!LHOeC4*3t${ zF;Q5z&SHTm3c_3Tuh!%3t#>>lrSgaJ}Q$;tG1{V z6~*YhY9q`(+N;_`L7I)Gg)lAJfr_wtk648ZToEZcEkp>VxC=%IuezYIKyA2ucdPFO z?gX;0eO&!X;9+3R-&^p%n)koD$}9d8Kl88qPbsd`Irg*&l1@#*&nUn(EcDZVgnEIx zO35@&$;>s6Z0Q8CzsmB=9+)953u&=3$8#`lB|V#LB|SUoIb<8@IZ4ljFaSz7=~aPP zdo)1sUTFJB+sXT(T}|39z6RR0CDE^gc0FlVUG?w{(A$MQ12|s3QTFix;Bk{$_qr&& zoe`2#^B1O~31LVRvRlq%x;q)unCzJ@k;DZlGJ8vPOGp_Z;x7I6gZo(ic~E;h4rx}5 zCvQY&gx?}_C@MFo)(bafF`7;#2>v25VJbQg`;=CGKBzk8;>ifvw6usl0b|jd+SnFk zqDT>O4MJllWcJYTXciGgTf~>GvLmUaAiAIl^5KSy21Ozt1U_neMK*Msan^wAYM?@h zRdF3(VU?B1MUxH#J3xf@g@Zbt{tFx#a0Okhj2gJoM<`HnV--#dg{%%)>0 zO`G=F>F#*mZN;?e_qjVWMvo%Tx0P#D-(^s#FcYLE5L~42hF+MZE3X@4D(%+yZiT%V zoddIG4*2wYm9M~m91#(MVULJR7VEHT!V0-7B1FX?R4OaQZwab%3{CQD zqL>oFW{5LVN|aT0DiKeM?MT{*6jZAlVZN zg*2;D6^W!%k{pRA<8maTcI>1-OhiP1{5$LdcE?wJ&w0n0?JYEQ=bO5JJo(x6Pp_}} zh6=vnyl?n%yW)Fuk=yXs5Yu$J5ImR<9{e%?*;}8!wdOxj@Sn~5&pxhJ{O1>~MVIgW z{`dN`{vQl(R0ng-$Cc{AoO@6ou7A*UziG|YN#?F7)opyH>WQCE)0Mxwz(gcC3Jtsb&Oi>q#3 zme6o$Zw#UGB9$)g1p+jm5<&xYww-(BE38)PHu180g!Hr%8sZMLXg7%Jo`o}`jTwc3 z?PaUml4dGX0i;d9Oj{;OfTzp~hrQC7VFsAKeo$fXR6%OYLfQbkAk_p8LQ=3gT-2yZ3F&vz_vRU6wRHl$D^e!V10`j&vU26Z z)R?sk?Qf~*`{`FuE3Am`Knz1Q^XLEZ4@=MgY2~FQr@&1oZis}u+=Hmraa{|D zmf|rXB1V%lLgZb6EZx5YyerTWk$VreU7(8Inu51I?`>Z?wC>%#$QFII1z$(r*RgbR z-M1Go=WfWl75DB%y6E!$%GI3Z)?Lk?*VJ!XD0j;jHT8F;?4b|){_=YtO)a%oIcUAN-x$=Q2; zPYjAjk)1Dao(R%j0U{#+;W!LHuwJ^fNv_Kot^WeF9VKWh2OuVs%#`JG0`&(eqZOPI zFvg)#3;?zjKxZYXE}>DE!n6hG45=<}F=jXtX_lwPPw|u)S38wUhvPi;DOUjBfbF*J(^LpYxW(&-1Xj`m>QwN0ho#iuZKRUW(z+5b4GUAAlI6PI8u=jS=`#OpgUv&1nY7 zJ{grjN6bkYw^7-UP~)XIK}rdA1PRo+H!)C$dWqX4*A>7E_E!Nfc-8P?SPfW*b^!X0 zr{F!qCS|;orLBw{p3?O-$dX|}U}340CEF~;M56Ws2?S9CaM8r-4Q+rZ*EeyLJPICd z&z@4KZA*(nEIuWx&6Qt&X__oV`GdFbt-d>uy{XjfTkcT2-7B@9HGkT?`d#JF`E~F4 zoc;X&)?H{N$SXs6{@o)$Jl0EJ4p`uFi09;-Ur|tHPm@deBtS$8BAL92uc5I32@l}0 zHnXK3MHX!jV>fMUA6KJVco<+yHjhZk_eI-%H=$+@xj<$)ug23}uJu_*ds{kRgukHV{?-)~7%R4-~ zop)@1_Lr`4@~;2$8aH3HeGTt9R^lf>*h@lRu^IR-UWDp(4Z|d+#yDStX9!!ZSP^)( zVXG6X4yeF^Amn8=5W=Z4kgWvXs%1ui=zB7iN|=IojN8K&)iRZk#4ETA^@xNr29Aa1 zRY35Unlo^}L#ED5^Ez(vF;olG7qyMqy$=rDKd^LT<**VsrquQ=IyPM1+qQeIzi{P( zUCV;v>n*sB{L*!#xT|S#_;YXdMpN^H{tx<>W^&!f*Y^zmT>QnYN4M6R_>IP<2Z0X) zOJ|lx);kA&{s+J4f7HL$I8v-{-RRi2-0(@$!=~JkHy_JN_@dHrX`{X=d*;Ep`{$Nj zD;JdDx0L$+3PWF4*-&@bO1dM=lvz0v_N(?tBnf#DNH91fk-1bX19{g7rvKm0L=&YM zZzM7u7bQ6nPYTHtOv8Zba|e7NAgVb8{3yQ@@i2BzV)X`Ae~8r>RH}=77nu{}*;Gut zh9h;62n0=Xd@2%^WidXPks%KpA)1G14lRv}aju9Jnnh_Esz0Tk)V3&WOR+kru)$(O zU}1PMqp*Qu=e`^Z8Fz)FCc;hXJV#G5@5&d)@Gh{N&BH9NeOy4q#@|txj+wOCFtF&F>-l4 zKrXq-X)>#6XtUVm1>4F;U7{ASemIg9kyQ#Gb)% zW#-ZlCy=FaeLYcbs&y8!pb0@jDFUxqPdd0I$bv?6p?eo5iHy7r4ag#}-U3^nXX^`W zFwX{;F0R1$sNfp=tta;S+^)k#yZ`pFLe1`c&2FV;Z_)0(U3ahPc2l-})5dzO3uiW+ zl(VkjXv{krvz-M;>zbo=>F9>t^S&M+YFy zKnJoiWF6>0Zr>ZEtDy&QViW^C-@p(_@tJ2!WDq=+Mk~lLZTk9&Oxaepd{JQFb20pt zW!U6*VDf5Pxyz-+`ZmCIy7r z0WulLI|7gbcZAj)p{3V1s(0NxaOXgFTB+W>XhW%N`@+*v@Vt`uyt2};+PHq`l;SzP zaHi<4dH>vd=d!k?YQ^2T@LJJTy?8vwHa>U23fSx$33>EE!_a>2(SFxZzavbKgsaei zPG{hgJbp0;2?j%F{3ifG67dzHqL613@kyfOH6CzD0Q0=&uV4aFV+SHs#3$Nj1XcU1 znvp&({xQs=^hiTcZCYrW{x#M9Z&cf_sa^j{RsDvlerCnb^qY2yaXn{F(KL7gSc73> UT$^Oz`FREm!t;8NWFo=;13$YOJ^%m! literal 0 HcmV?d00001 diff --git a/scripts_4/adapter.py b/scripts_4/adapter.py new file mode 100644 index 0000000..80d0834 --- /dev/null +++ b/scripts_4/adapter.py @@ -0,0 +1,101 @@ +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +import pickle + +import numpy as np + + +def _safe_float(value: float) -> float: + return float(np.asarray(value, dtype=np.float64).reshape(-1)[0]) + + +@dataclass +class FrequencyScaleAdapter: + frequencies_hz: np.ndarray + scale_values: np.ndarray + shrinkage: float = 0.2 + extrapolation_decay_hz: float = 0.25 + min_scale: float = 0.4 + max_scale: float = 2.5 + + def __post_init__(self) -> None: + freqs = np.asarray(self.frequencies_hz, dtype=np.float64).reshape(-1) + scales = np.asarray(self.scale_values, dtype=np.float64).reshape(-1) + if freqs.size == 0: + raise ValueError("frequencies_hz cannot be empty.") + if freqs.size != scales.size: + raise ValueError("frequencies_hz and scale_values must have the same length.") + order = np.argsort(freqs) + freqs = freqs[order] + scales = np.clip(scales[order], self.min_scale, self.max_scale) + # Pull scales toward 1.0 to reduce overfitting with tiny TMD samples. + scales = 1.0 + (scales - 1.0) * (1.0 - float(np.clip(self.shrinkage, 0.0, 0.95))) + self.frequencies_hz = freqs + self.scale_values = scales + + def _inside_range_weight(self, frequency_hz: float) -> float: + left = self.frequencies_hz[0] + right = self.frequencies_hz[-1] + f = _safe_float(frequency_hz) + if left <= f <= right: + return 1.0 + distance = min(abs(f - left), abs(f - right)) + decay = max(_safe_float(self.extrapolation_decay_hz), 1e-6) + return float(np.exp(-distance / decay)) + + def scale_at(self, frequency_hz: float) -> float: + f = _safe_float(frequency_hz) + boundary_interp = float(np.interp(f, self.frequencies_hz, self.scale_values)) + in_range_weight = self._inside_range_weight(f) + scale = 1.0 + in_range_weight * (boundary_interp - 1.0) + return float(np.clip(scale, self.min_scale, self.max_scale)) + + def predict(self, base_rms: float, frequency_hz: float) -> float: + return max(_safe_float(base_rms), 0.0) * self.scale_at(frequency_hz) + + def to_payload(self) -> dict: + return { + "frequencies_hz": self.frequencies_hz.tolist(), + "scale_values": self.scale_values.tolist(), + # scale_values are already post-shrinkage values after __post_init__. + "shrinkage": 0.0, + "extrapolation_decay_hz": float(self.extrapolation_decay_hz), + "min_scale": float(self.min_scale), + "max_scale": float(self.max_scale), + "already_shrunk": True, + } + + @classmethod + def from_payload(cls, payload: dict) -> "FrequencyScaleAdapter": + shrinkage = float(payload.get("shrinkage", 0.2)) + if bool(payload.get("already_shrunk", False)): + shrinkage = 0.0 + return cls( + frequencies_hz=np.asarray(payload["frequencies_hz"], dtype=np.float64), + scale_values=np.asarray(payload["scale_values"], dtype=np.float64), + shrinkage=shrinkage, + extrapolation_decay_hz=float(payload.get("extrapolation_decay_hz", 0.25)), + min_scale=float(payload.get("min_scale", 0.4)), + max_scale=float(payload.get("max_scale", 2.5)), + ) + + +def save_adapter(adapter: FrequencyScaleAdapter, output_path: Path, extra: dict | None = None) -> None: + output_path.parent.mkdir(parents=True, exist_ok=True) + payload = { + "adapter_type": "frequency_scale_interp_v1", + "adapter": adapter.to_payload(), + "extra": extra or {}, + } + with output_path.open("wb") as handle: + pickle.dump(payload, handle) + + +def load_adapter(adapter_path: Path) -> tuple[FrequencyScaleAdapter, dict]: + with adapter_path.open("rb") as handle: + payload = pickle.load(handle) + adapter = FrequencyScaleAdapter.from_payload(payload["adapter"]) + extra = payload.get("extra", {}) + return adapter, extra diff --git a/scripts_4/predict_tmd_single.py b/scripts_4/predict_tmd_single.py new file mode 100644 index 0000000..441f9c2 --- /dev/null +++ b/scripts_4/predict_tmd_single.py @@ -0,0 +1,78 @@ +from __future__ import annotations + +import argparse +import importlib.util +import pickle +from pathlib import Path +import sys + +PROJECT_ROOT = Path(__file__).resolve().parents[1] + +try: + from .adapter import load_adapter +except ImportError: + from adapter import load_adapter + + +def _load_module(module_name: str, file_path: Path): + spec = importlib.util.spec_from_file_location(module_name, file_path) + if spec is None or spec.loader is None: + raise ImportError(f"Cannot load module from {file_path}") + module = importlib.util.module_from_spec(spec) + sys.modules[module_name] = module + spec.loader.exec_module(module) + return module + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Predict TMD RMS using baseline model + frequency adapter.") + parser.add_argument("--file", type=str, required=True) + parser.add_argument("--checkpoint", type=str, default="checkpoints_final/task1_final/task1_final_model.pkl") + parser.add_argument("--adapter", type=str, default="checkpoints_final/task1_final/task1_tmd_adapter.pkl") + parser.add_argument("--data-dir", type=str, default="downloads/TMD") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + config_module = _load_module("config", PROJECT_ROOT / "scripts" / "config.py") + dataset_module = _load_module("task1_dataset", PROJECT_ROOT / "scripts" / "dataset.py") + make_experiment_config = getattr(config_module, "make_experiment_config") + build_record_from_file = getattr(dataset_module, "build_record_from_file") + + config = make_experiment_config() + config.data.data_dir = args.data_dir + config.data.__post_init__() + + ckpt_path = (PROJECT_ROOT / args.checkpoint).resolve() + adapter_path = (PROJECT_ROOT / args.adapter).resolve() + with ckpt_path.open("rb") as handle: + payload = pickle.load(handle) + model = payload["model"] + adapter, extra = load_adapter(adapter_path) + + file_path = Path(args.file).resolve() + record = build_record_from_file(file_path, config.data) + + pred_tr = max(float(model.predict([record.features.tolist()])[0]), config.data.normalization_eps) + pred_base = pred_tr * record.x_rms + pred_tmd = adapter.predict(pred_base, record.frequency_hz) + + true_rms = float(record.y_rms) + base_err = abs(pred_base - true_rms) / max(abs(true_rms), 1e-12) * 100.0 + tmd_err = abs(pred_tmd - true_rms) / max(abs(true_rms), 1e-12) * 100.0 + + print(f"Checkpoint: {ckpt_path}") + print(f"Adapter: {adapter_path}") + print(f"Adapter extra: {extra}") + print(f"Input file: {file_path}") + print(f"Frequency (Hz): {record.frequency_hz:.6f}") + print(f"x_rms: {record.x_rms:.6f}") + print(f"True y_rms: {true_rms:.6f}") + print(f"Base pred y_rms: {pred_base:.6f} | error={base_err:.4f}%") + print(f"Adapted pred y_rms: {pred_tmd:.6f} | error={tmd_err:.4f}%") + print(f"Applied scale k(f): {adapter.scale_at(record.frequency_hz):.6f}") + + +if __name__ == "__main__": + main() diff --git a/scripts_4/train_tmd_adapter.py b/scripts_4/train_tmd_adapter.py new file mode 100644 index 0000000..b62a871 --- /dev/null +++ b/scripts_4/train_tmd_adapter.py @@ -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()