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:
78
scripts_4/predict_tmd_single.py
Normal file
78
scripts_4/predict_tmd_single.py
Normal file
@@ -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()
|
||||
Reference in New Issue
Block a user