from __future__ import annotations from dataclasses import dataclass, field from pathlib import Path import re from typing import Any import numpy as np import pandas as pd import torch from torch.utils.data import DataLoader, Dataset try: from .config import CORE_FEATURE_NAMES, DataConfig, ExperimentConfig except ImportError: from config import CORE_FEATURE_NAMES, DataConfig, ExperimentConfig def calculate_rms(signal: np.ndarray) -> float: signal = np.asarray(signal, dtype=np.float64).reshape(-1) return float(np.sqrt(np.mean(np.square(signal)))) def extract_frequency_hz(file_name: str) -> float | None: match = re.search(r"(\d+(?:\.\d+)?)Hz", file_name, flags=re.IGNORECASE) if match is None: return None return float(match.group(1)) def estimate_sampling_rate(time_values: np.ndarray) -> float: if time_values.size < 2: return 100.0 dt = np.diff(time_values) dt = dt[np.isfinite(dt)] dt = dt[dt > 0.0] if dt.size == 0: return 100.0 return float(1.0 / np.median(dt)) def get_middle_segment( time_values: np.ndarray, x_values: np.ndarray, y_values: np.ndarray, config: DataConfig, ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: length = len(x_values) search_start = int(length * config.middle_segment_start_ratio) search_end = int(length * config.middle_segment_end_ratio) search_start = max(0, min(search_start, length - 1)) search_end = max(search_start + 1, min(search_end, length)) if search_end - search_start < config.min_segment_length: center = length // 2 half = config.min_segment_length // 2 search_start = max(0, center - half) search_end = min(length, search_start + config.min_segment_length) search_start = max(0, search_end - config.min_segment_length) search_length = search_end - search_start window_length = max(config.min_segment_length, int(length * config.steady_window_ratio)) window_length = min(window_length, search_length) if search_length <= window_length: return ( time_values[search_start:search_end], x_values[search_start:search_end], y_values[search_start:search_end], ) stride = max(1, int(length * config.steady_window_stride_ratio)) candidate_y = y_values[search_start:search_end] candidate_rms = calculate_rms(candidate_y) best_score = float("inf") best_slice = slice(search_start, search_start + window_length) for window_start in range(search_start, search_end - window_length + 1, stride): window_end = window_start + window_length window_y = y_values[window_start:window_end] split_windows = np.array_split(window_y, config.stability_subwindow_count) split_rms = np.asarray([calculate_rms(chunk) for chunk in split_windows], dtype=np.float64) split_peaks = np.asarray([float(np.max(np.abs(chunk))) for chunk in split_windows], dtype=np.float64) rms_cv = float(split_rms.std() / max(split_rms.mean(), config.normalization_eps)) peak_cv = float(split_peaks.std() / max(split_peaks.mean(), config.normalization_eps)) window_rms = calculate_rms(window_y) # Prefer windows that are stable inside the middle candidate region while keeping enough energy. score = rms_cv + 0.35 * peak_cv - 0.05 * (window_rms / max(candidate_rms, config.normalization_eps)) if score < best_score: best_score = score best_slice = slice(window_start, window_end) return time_values[best_slice], x_values[best_slice], y_values[best_slice] def _compute_windowed_spectrum(signal: np.ndarray, sampling_rate: float) -> tuple[np.ndarray, np.ndarray]: signal = np.asarray(signal, dtype=np.float64).reshape(-1) if signal.size < 4: return np.asarray([], dtype=np.float64), np.asarray([], dtype=np.float64) centered = signal - np.mean(signal) window = np.hanning(signal.size) scale = max(np.sum(window), 1e-12) fft_values = np.fft.rfft(centered * window) freqs = np.fft.rfftfreq(signal.size, d=1.0 / sampling_rate) magnitudes = (2.0 / scale) * np.abs(fft_values) return freqs, magnitudes def _parabolic_peak_frequency(freqs: np.ndarray, magnitudes: np.ndarray, peak_index: int) -> float: if peak_index <= 0 or peak_index >= magnitudes.size - 1: return float(freqs[peak_index]) alpha = magnitudes[peak_index - 1] beta = magnitudes[peak_index] gamma = magnitudes[peak_index + 1] denominator = alpha - 2.0 * beta + gamma if abs(denominator) < 1e-12: return float(freqs[peak_index]) offset = 0.5 * (alpha - gamma) / denominator bin_width = float(freqs[1] - freqs[0]) return float(freqs[peak_index] + offset * bin_width) def get_dominant_frequency(signal: np.ndarray, sampling_rate: float) -> float: freqs, magnitudes = _compute_windowed_spectrum(signal, sampling_rate) if magnitudes.size == 0: return 0.0 magnitudes[0] = 0.0 band_mask = (freqs >= 0.1) & (freqs <= 5.0) if not np.any(band_mask): return 0.0 band_magnitudes = np.where(band_mask, magnitudes, 0.0) peak_index = int(np.argmax(band_magnitudes)) return _parabolic_peak_frequency(freqs, magnitudes, peak_index) def compute_harmonic_fit_features(signal: np.ndarray, sampling_rate: float, frequency_hz: float) -> tuple[float, float]: signal = np.asarray(signal, dtype=np.float64).reshape(-1) if signal.size < 4 or frequency_hz <= 0.0: return 0.0, 1.0 time_axis = np.arange(signal.size, dtype=np.float64) / max(sampling_rate, 1e-12) omega_t = 2.0 * np.pi * frequency_hz * time_axis design = np.stack([np.sin(omega_t), np.cos(omega_t), np.ones_like(omega_t)], axis=1) coefficients, _, _, _ = np.linalg.lstsq(design, signal, rcond=None) fitted = design @ coefficients harmonic_amplitude = float(np.sqrt(coefficients[0] ** 2 + coefficients[1] ** 2)) residual = signal - fitted residual_ratio = calculate_rms(residual) / max(calculate_rms(signal), 1e-6) return harmonic_amplitude, float(residual_ratio) def compute_spectral_features(signal: np.ndarray, sampling_rate: float) -> tuple[float, float, float, float, float]: freqs, amplitudes = _compute_windowed_spectrum(signal, sampling_rate) if amplitudes.size == 0: return 0.0, 0.0, 0.0, 0.0, 0.0 powers = np.square(amplitudes) amplitudes[0] = 0.0 powers[0] = 0.0 band_mask = (freqs >= 0.1) & (freqs <= 5.0) if not np.any(band_mask): return 0.0, 0.0, 0.0, 0.0, 0.0 band_amplitudes = np.where(band_mask, amplitudes, 0.0) dominant_index = int(np.argmax(band_amplitudes)) dominant_amplitude = float(amplitudes[dominant_index]) dominant_freq = _parabolic_peak_frequency(freqs, amplitudes, dominant_index) local_mask = np.abs(freqs - dominant_freq) <= 0.10 dominant_energy = float(powers[local_mask].sum()) total_band_energy = float(powers[band_mask].sum()) dominant_energy_ratio = dominant_energy / max(total_band_energy, 1e-12) spectral_centroid = float((freqs[band_mask] * powers[band_mask]).sum() / max(total_band_energy, 1e-12)) background_mask = band_mask & (~local_mask) background_level = float(np.median(amplitudes[background_mask])) if np.any(background_mask) else 0.0 spectral_peak_prominence = dominant_amplitude / max(background_level, 1e-6) half_power_level = dominant_amplitude / np.sqrt(2.0) left_index = dominant_index right_index = dominant_index while left_index > 0 and amplitudes[left_index] >= half_power_level: left_index -= 1 while right_index < amplitudes.size - 1 and amplitudes[right_index] >= half_power_level: right_index += 1 half_power_bandwidth = float(freqs[right_index] - freqs[left_index]) if right_index > left_index else 0.0 return dominant_amplitude, dominant_energy_ratio, spectral_centroid, spectral_peak_prominence, half_power_bandwidth def extract_core_features( signal: np.ndarray, sampling_rate: float, original_length: int, config: DataConfig, ) -> np.ndarray: x_rms = calculate_rms(signal) x_peak = float(np.max(np.abs(signal))) if signal.size > 0 else 0.0 x_peak_to_peak = float(np.max(signal) - np.min(signal)) if signal.size > 0 else 0.0 crest_factor = x_peak / max(x_rms, config.normalization_eps) dominant_frequency = get_dominant_frequency(signal, sampling_rate) frequency_squared = dominant_frequency * dominant_frequency inverse_frequency = 1.0 / max(dominant_frequency, 1e-6) log_frequency = float(np.log(max(dominant_frequency, 1e-6))) middle_length_ratio = float(signal.size) / float(max(original_length, 1)) ( dominant_amplitude, dominant_energy_ratio, spectral_centroid, spectral_peak_prominence, half_power_bandwidth, ) = compute_spectral_features(signal, sampling_rate) harmonic_fit_amplitude, harmonic_fit_residual_ratio = compute_harmonic_fit_features( signal, sampling_rate, dominant_frequency, ) signal_mean = float(np.mean(signal)) if signal.size > 0 else 0.0 return np.asarray( [ dominant_frequency, frequency_squared, inverse_frequency, log_frequency, x_rms, x_peak, x_peak_to_peak, crest_factor, middle_length_ratio, dominant_amplitude, dominant_energy_ratio, harmonic_fit_amplitude, harmonic_fit_residual_ratio, spectral_peak_prominence, half_power_bandwidth, spectral_centroid, signal_mean, ], dtype=np.float32, ) def _value_frame(df: pd.DataFrame, config: DataConfig, sensor_code: str, value_column: str) -> pd.DataFrame: sensor_df = df.loc[df[config.code_column] == sensor_code, [config.time_column, value_column]].copy() sensor_df = sensor_df.sort_values(config.time_column) sensor_df = sensor_df.drop_duplicates(subset=config.time_column, keep="first") sensor_df[config.time_column] = sensor_df[config.time_column].astype("float64") sensor_df[value_column] = sensor_df[value_column].astype("float32") return sensor_df def load_aligned_signals(file_path: Path, config: DataConfig) -> tuple[np.ndarray, np.ndarray, np.ndarray, int]: df = pd.read_csv(file_path) base_df = _value_frame(df, config, config.base_sensor_code, config.base_axis) response_df = _value_frame(df, config, config.response_sensor_code, config.response_axis) if base_df.empty or response_df.empty: raise ValueError(f"Missing required sensor in {file_path.name}") aligned = base_df.rename(columns={config.base_axis: "base_signal"}).merge( response_df.rename(columns={config.response_axis: "response_signal"}), on=config.time_column, how="left", ) interpolation_count = int(aligned["response_signal"].isna().sum()) aligned["response_signal"] = aligned["response_signal"].interpolate( method=config.interpolation_method, limit_direction="both", ).ffill().bfill() if aligned["response_signal"].isna().any(): raise ValueError(f"Remaining NaN after interpolation in {file_path.name}") time_values = aligned[config.time_column].to_numpy(dtype=np.float64) x_values = aligned["base_signal"].to_numpy(dtype=np.float32) y_values = aligned["response_signal"].to_numpy(dtype=np.float32) return time_values, x_values, y_values, interpolation_count @dataclass class SampleRecord: file_path: Path split: str frequency_hz: float features: torch.Tensor x_rms: torch.Tensor y_rms: torch.Tensor target_y_rms: torch.Tensor target_log_y_rms: torch.Tensor time_middle: torch.Tensor x_middle: torch.Tensor y_middle: torch.Tensor sampling_rate: float interpolation_count: int = 0 @dataclass class SplitLoadReport: split: str loaded_files: list[str] = field(default_factory=list) skipped_files: list[tuple[str, str]] = field(default_factory=list) interpolated_files: dict[str, int] = field(default_factory=dict) def to_lines(self) -> list[str]: lines = [f"[{self.split}] loaded={len(self.loaded_files)} skipped={len(self.skipped_files)}"] for file_name, reason in self.skipped_files: lines.append(f" - skipped {file_name}: {reason}") for file_name, count in self.interpolated_files.items(): if count > 0: lines.append(f" - interpolated {file_name}: missing_points={count}") return lines @dataclass class NormalizationStats: feature_mean: torch.Tensor feature_std: torch.Tensor target_mean: torch.Tensor target_std: torch.Tensor class FeatureDataset(Dataset): def __init__(self, records: list[SampleRecord], normalization: NormalizationStats) -> None: self.records = records self.normalization = normalization def __len__(self) -> int: return len(self.records) def __getitem__(self, index: int) -> dict[str, Any]: record = self.records[index] feature_norm = (record.features - self.normalization.feature_mean) / self.normalization.feature_std target_norm = (record.target_y_rms - self.normalization.target_mean) / self.normalization.target_std return { "x": feature_norm, "target": target_norm, "target_raw": record.target_y_rms, "target_log_raw": record.target_log_y_rms, "x_rms_raw": record.x_rms, "y_rms_raw": record.y_rms, "frequency_hz": torch.tensor(record.frequency_hz, dtype=torch.float32), "features_raw": record.features, "file_name": record.file_path.name, } def list_split_files(config: DataConfig, split: str) -> list[Path]: split_dir = getattr(config, f"{split}_dir") return sorted(split_dir.glob(config.csv_pattern)) def build_record_from_file(file_path: Path, split: str, config: DataConfig) -> SampleRecord: time_values, x_values, y_values, interpolation_count = load_aligned_signals(file_path, config) time_middle, x_middle, y_middle = get_middle_segment(time_values, x_values, y_values, config) sampling_rate = estimate_sampling_rate(time_middle) frequency_hz = extract_frequency_hz(file_path.name) if frequency_hz is None: frequency_hz = get_dominant_frequency(x_middle, sampling_rate) x_rms = calculate_rms(x_middle) y_rms = calculate_rms(y_middle) transmission_ratio = float(y_rms / max(x_rms, config.normalization_eps)) target_y_rms = transmission_ratio target_log_y_rms = float(np.log(max(transmission_ratio, config.normalization_eps))) features = extract_core_features(x_middle, sampling_rate, len(x_values), config) return SampleRecord( file_path=file_path, split=split, frequency_hz=frequency_hz, features=torch.tensor(features, dtype=torch.float32), x_rms=torch.tensor([x_rms], dtype=torch.float32), y_rms=torch.tensor([y_rms], dtype=torch.float32), target_y_rms=torch.tensor([target_y_rms], dtype=torch.float32), target_log_y_rms=torch.tensor([target_log_y_rms], dtype=torch.float32), time_middle=torch.tensor(time_middle, dtype=torch.float64), x_middle=torch.tensor(x_middle[:, None], dtype=torch.float32), y_middle=torch.tensor(y_middle[:, None], dtype=torch.float32), sampling_rate=sampling_rate, interpolation_count=interpolation_count, ) def load_split_records(config: DataConfig, split: str) -> tuple[list[SampleRecord], SplitLoadReport]: records: list[SampleRecord] = [] report = SplitLoadReport(split=split) for file_path in list_split_files(config, split): try: record = build_record_from_file(file_path, split, config) except ValueError as error: report.skipped_files.append((file_path.name, str(error))) continue records.append(record) report.loaded_files.append(file_path.name) if record.interpolation_count > 0: report.interpolated_files[file_path.name] = record.interpolation_count if not records: raise RuntimeError(f"No usable records found for split='{split}'.") return records, report def fit_normalization(records: list[SampleRecord], config: DataConfig) -> NormalizationStats: feature_all = torch.stack([record.features for record in records], dim=0) target_all = torch.cat([record.target_y_rms for record in records], dim=0) feature_std = torch.clamp(feature_all.std(dim=0, unbiased=False), min=config.normalization_eps) target_std = torch.clamp(target_all.std(dim=0, unbiased=False).view(1), min=config.normalization_eps) return NormalizationStats( feature_mean=feature_all.mean(dim=0), feature_std=feature_std, target_mean=target_all.mean(dim=0, keepdim=True), target_std=target_std, ) def build_datasets( config: ExperimentConfig | DataConfig, ) -> tuple[dict[str, FeatureDataset], dict[str, list[SampleRecord]], dict[str, SplitLoadReport]]: data_config = config.data if isinstance(config, ExperimentConfig) else config train_records, train_report = load_split_records(data_config, "train") normalization = fit_normalization(train_records, data_config) val_records, val_report = load_split_records(data_config, "val") test_records, test_report = load_split_records(data_config, "test") raw_records = { "train": train_records, "val": val_records, "test": test_records, } datasets = { split: FeatureDataset(records, normalization) for split, records in raw_records.items() } reports = { "train": train_report, "val": val_report, "test": test_report, } return datasets, raw_records, reports def build_dataloaders( config: ExperimentConfig | DataConfig, ) -> tuple[dict[str, DataLoader], dict[str, FeatureDataset], dict[str, list[SampleRecord]], dict[str, SplitLoadReport]]: experiment_config = config if isinstance(config, ExperimentConfig) else ExperimentConfig(data=config) datasets, raw_records, reports = build_datasets(experiment_config) batch_size = experiment_config.train.batch_size loaders = { "train": DataLoader(datasets["train"], batch_size=batch_size, shuffle=True), "val": DataLoader(datasets["val"], batch_size=batch_size, shuffle=False), "test": DataLoader(datasets["test"], batch_size=batch_size, shuffle=False), } return loaders, datasets, raw_records, reports def report_to_text(reports: dict[str, SplitLoadReport]) -> str: lines: list[str] = [] for split in ("train", "val", "test"): lines.extend(reports[split].to_lines()) return "\n".join(lines) def checkpoint_normalization_payload(normalization: NormalizationStats) -> dict[str, list[float]]: return { "feature_mean": normalization.feature_mean.tolist(), "feature_std": normalization.feature_std.tolist(), "target_mean": normalization.target_mean.tolist(), "target_std": normalization.target_std.tolist(), "feature_names": list(CORE_FEATURE_NAMES), } def normalization_from_payload(payload: dict[str, Any]) -> NormalizationStats: return NormalizationStats( feature_mean=torch.tensor(payload["feature_mean"], dtype=torch.float32), feature_std=torch.tensor(payload["feature_std"], dtype=torch.float32), target_mean=torch.tensor(payload["target_mean"], dtype=torch.float32), target_std=torch.tensor(payload["target_std"], dtype=torch.float32), )