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