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:
CrbnsCat10n
2026-05-06 20:57:43 +08:00
parent f3766c74d6
commit 79f3de05f7
9 changed files with 390 additions and 0 deletions

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