modified: Figure_1.png
modified: checkpoints/forward/best_tcn_model.pt modified: checkpoints/forward/training_history.csv new file: checkpoints_mlp/task1_feature_mlp/best_feature_mlp.pt new file: checkpoints_mlp/task1_feature_mlp/training_history.csv new file: checkpoints_rms/forward_rms/best_rms_model.pt new file: checkpoints_rms/forward_rms/training_history.csv modified: evaluation_outputs/forward/evaluation_forward_test_b0_s0.png new file: evaluation_outputs/forward/evaluation_forward_val_b0_s0.png new file: evaluation_outputs/forward/evaluation_forward_val_b3_s0.png new file: evaluation_outputs/forward_rms/evaluation_test_all_samples.csv new file: evaluation_outputs/forward_rms/evaluation_test_s0.png new file: evaluation_outputs/forward_rms/evaluation_train_all_samples.csv new file: evaluation_outputs/forward_rms/evaluation_val_all_samples.csv new file: evaluation_outputs/forward_rms/evaluation_val_s0.png new file: evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png new file: evaluation_outputs/task1_feature_mlp/evaluation_train_all_samples.csv new file: evaluation_outputs/task1_feature_mlp/evaluation_train_curve.png new file: evaluation_outputs/task1_feature_mlp/evaluation_val_all_samples.csv new file: evaluation_outputs/task1_feature_mlp/evaluation_val_curve.png new file: evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png new file: evaluation_outputs/task1_feature_mlp/harmonic_5mm_0.75Hz_prediction.png new file: evaluation_outputs/task1_feature_mlp/harmonic_5mm_1.55Hz_prediction.png new file: scripts/__pycache__/config.cpython-310.pyc new file: scripts/__pycache__/dataset.cpython-310.pyc new file: scripts/__pycache__/model.cpython-310.pyc new file: scripts/config.py new file: scripts/dataset.py new file: scripts/evaluate.py new file: scripts/model.py new file: scripts/predict_single.py new file: scripts/train.py modified: src/__pycache__/config.cpython-310.pyc modified: src/__pycache__/dataset.cpython-310.pyc modified: src/__pycache__/model.cpython-310.pyc modified: src/config.py modified: src/dataset.py modified: src/model.py new file: src_new/__pycache__/config.cpython-310.pyc new file: src_new/__pycache__/dataset.cpython-310.pyc new file: src_new/__pycache__/evaluate.cpython-310.pyc new file: src_new/__pycache__/model.cpython-310.pyc new file: src_new/__pycache__/train.cpython-310.pyc new file: src_new/config.py new file: src_new/dataset.py new file: src_new/evaluate.py new file: src_new/model.py new file: src_new/train.py
This commit is contained in:
478
scripts/dataset.py
Normal file
478
scripts/dataset.py
Normal file
@@ -0,0 +1,478 @@
|
||||
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),
|
||||
)
|
||||
Reference in New Issue
Block a user