new file: Figure_1.png
deleted: best_model.pth new file: checkpoints/forward/best_tcn_model.pt new file: checkpoints/forward/training_history.csv new file: evaluation_outputs/forward/evaluation_forward_test_b0_s0.png new file: sanity_check_alignment_forward.png modified: src/__pycache__/config.cpython-310.pyc modified: src/__pycache__/config.cpython-314.pyc modified: src/__pycache__/dataset.cpython-310.pyc modified: src/__pycache__/dataset.cpython-314.pyc modified: src/__pycache__/model.cpython-310.pyc modified: src/__pycache__/model.cpython-314.pyc modified: src/config.py modified: src/dataset.py modified: src/evaluate.py modified: src/model.py new file: src/sanity_check.py modified: src/train.py renamed: src/__init__.py -> src_old/__init__.py new file: src_old/__pycache__/config.cpython-310.pyc new file: src_old/__pycache__/config.cpython-314.pyc new file: src_old/__pycache__/dataset.cpython-310.pyc new file: src_old/__pycache__/dataset.cpython-314.pyc new file: src_old/__pycache__/evaluate.cpython-314.pyc new file: src_old/__pycache__/model.cpython-310.pyc new file: src_old/__pycache__/model.cpython-314.pyc new file: src_old/__pycache__/train.cpython-310.pyc new file: src_old/__pycache__/train.cpython-314.pyc new file: src_old/config.py new file: src_old/dataset.py new file: src_old/evaluate.py new file: src_old/model.py new file: src_old/train.py deleted: test_results.png
This commit is contained in:
BIN
Figure_1.png
Normal file
BIN
Figure_1.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 190 KiB |
BIN
best_model.pth
BIN
best_model.pth
Binary file not shown.
BIN
checkpoints/forward/best_tcn_model.pt
Normal file
BIN
checkpoints/forward/best_tcn_model.pt
Normal file
Binary file not shown.
14
checkpoints/forward/training_history.csv
Normal file
14
checkpoints/forward/training_history.csv
Normal file
@@ -0,0 +1,14 @@
|
|||||||
|
epoch,lr,train_loss,train_time_loss,train_weighted_time_loss,train_fft_loss,train_rms_loss,train_scale_loss,train_underestimate_loss,train_envelope_loss,train_rms_error,train_dominant_freq_error,val_loss,val_time_loss,val_weighted_time_loss,val_fft_loss,val_rms_loss,val_scale_loss,val_underestimate_loss,val_envelope_loss,val_rms_error,val_dominant_freq_error
|
||||||
|
1,0.001,32.23490004269582,0.9063257234838774,0.09063257336757093,2462.537229142099,0.5202015545570625,0.5072746667659508,0.23183272828189833,0.9878060196368199,0.5202015545570625,0.6667674588373308,16.848691019518622,0.5030307245665583,0.05030307315033058,712.1304384428879,0.21638367042459292,0.3054182426682834,0.16712227438030572,0.7986705179872184,0.21638367042459292,0.9571996228951049
|
||||||
|
2,0.001,21.43670712776904,0.9317702313639084,0.09317702508338217,1694.0047123747052,0.35049294160222105,0.3616391454102858,0.1121402242625097,0.672775794312639,0.35049294160222105,0.6002240721274537,18.15417513354071,0.5811605972462687,0.058116060392609956,948.0657793242356,0.23726162314414978,0.3922087324076685,0.15394838554142365,0.7320991713425209,0.23726162314414978,0.8361816405868994
|
||||||
|
3,0.001,18.081004331696708,0.949055330933265,0.09490553458344261,1372.9708672289578,0.28145609332143134,0.3131425976753235,0.09374795892750318,0.6080609613432074,0.28145609332143134,0.616063318727439,16.73266759412042,0.5420598767954727,0.05420598895128431,673.0913622625943,0.238979727286717,0.33492567107595245,0.1804515516449665,0.7293170896069757,0.238979727286717,0.7597824622933206
|
||||||
|
4,0.001,15.158426779621053,0.9718681652590914,0.09718681784030402,1005.0683164776497,0.2189214069325969,0.2836238658934269,0.08367646379695046,0.5845890998278024,0.2189214069325969,0.6179393338571753,17.278400026518725,0.5685602380283947,0.056856024560743366,892.716255977236,0.22618299260221678,0.36168746264844104,0.1538948353444194,0.6990852653980255,0.22618299260221678,0.7597824622933206
|
||||||
|
5,0.001,13.61913053044733,0.9698394334541177,0.09698394513776842,859.7200766509434,0.19106561557020782,0.2557676170232161,0.07612784807833861,0.5538543634257227,0.19106561557020782,0.6155450957215706,17.31632186626566,0.5785790409507423,0.057857905097048856,895.00206572434,0.21634554451909557,0.35666010605877846,0.1514351657302729,0.7220557270378902,0.21634554451909557,0.7402091190436203
|
||||||
|
6,0.001,12.678725962368947,0.9694307513956754,0.0969430767351164,788.4346546676924,0.1713324544845887,0.23463592369039105,0.07080332469195127,0.5323761788741598,0.1713324544845887,0.6164817462027068,18.423015594482422,0.5721222347226637,0.057212224808232535,888.629751271215,0.25913861959144985,0.4007708651238474,0.16395779438959113,0.7525067822686557,0.25913861959144985,0.7218985721265917
|
||||||
|
7,0.001,12.268796533908484,0.967802310889622,0.09678023287429,744.56112426182,0.1675024907684551,0.2299804154713199,0.06777841198029665,0.5231116253812358,0.1675024907684551,0.5961900787542253,17.48041952067408,0.5510032737049563,0.055100328578003524,704.0817605380354,0.2379121222886546,0.37086750232967836,0.19612120528673305,0.7370918125941835,0.2379121222886546,0.7088496766840734
|
||||||
|
8,0.001,11.889431161700555,0.967772865070487,0.09677728824317455,712.115890215028,0.1614244244289848,0.2226250537161557,0.06655820054968573,0.5118434342010966,0.1614244244289848,0.6022199021486582,18.410391971982758,0.5461424337378864,0.05461424415738418,738.1983037488214,0.2533088120920905,0.3795644687167529,0.21159803308546543,0.7770289305982918,0.2533088120920905,0.707797346401937
|
||||||
|
9,0.001,11.234980547203207,0.9654024828155086,0.09654024967326308,670.3140682004532,0.14827435323089924,0.20859187953876998,0.062428659900038874,0.4916352005499714,0.14827435323089924,0.5623616312479263,16.51186601046858,0.5469307437025267,0.05469307598882708,658.5074381335028,0.23887184570575581,0.3477302913008065,0.18191729235494958,0.6951427665250055,0.23887184570575581,0.7065345500594454
|
||||||
|
10,0.001,10.992490152143082,0.9626774315564137,0.09626774419591112,649.039663350807,0.14625250653557056,0.20260991113928128,0.061842711396374796,0.48383018318212256,0.14625250653557056,0.568747647623218,17.501161114922887,0.5506949491541961,0.05506949608439002,671.2075416301859,0.2541997463538729,0.373840541675173,0.19259302314884705,0.7480303932880533,0.2541997463538729,0.708428744563363
|
||||||
|
11,0.0005,10.806822551871246,0.9680146614335617,0.09680146787245318,637.6420590382702,0.14451587228280194,0.2035975716305229,0.0572100906404403,0.4776927213061531,0.14451587228280194,0.5574953460125556,17.13407661174906,0.5288547395632185,0.05288547495829648,618.1973874322299,0.2417878899081,0.34800343832065317,0.21005815672206468,0.7388678706925491,0.2417878899081,0.7122171335564683
|
||||||
|
12,0.0005,10.255595112746617,0.9616962041494981,0.09616962194723903,599.1373590433373,0.12593858316540718,0.18720266560338578,0.05598305191247249,0.4703599948365733,0.12593858316540718,0.5670544286673502,16.41022307297279,0.5580467443014013,0.055804675894564594,691.4893851444639,0.23402822017669678,0.3565609496215294,0.17255781141334567,0.6764166149599798,0.23402822017669678,0.6966426454535165
|
||||||
|
13,0.0005,9.979247138185322,0.9624017550135558,0.09624017732885648,580.7163422782467,0.12185739905063836,0.18553815514973873,0.05258095577657926,0.4598706980358879,0.12185739905063836,0.5524671217591512,16.273903156148975,0.5295915886245924,0.05295915985158805,541.1639246447332,0.23793373128463483,0.3361044392503541,0.20454522344315873,0.7062619534032099,0.23793373128463483,0.7056926858376642
|
||||||
|
BIN
evaluation_outputs/forward/evaluation_forward_test_b0_s0.png
Normal file
BIN
evaluation_outputs/forward/evaluation_forward_test_b0_s0.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 362 KiB |
BIN
sanity_check_alignment_forward.png
Normal file
BIN
sanity_check_alignment_forward.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 502 KiB |
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
149
src/config.py
149
src/config.py
@@ -1,33 +1,126 @@
|
|||||||
import os
|
from __future__ import annotations
|
||||||
|
|
||||||
# Data Configuration
|
from dataclasses import dataclass, field
|
||||||
# 使用基于当前文件的绝对路径拼接,以防止你在不同目录下运行报错
|
from pathlib import Path
|
||||||
DATA_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), 'downloads')
|
from typing import Literal
|
||||||
SEQ_LEN = 512 # 滑动窗口的长度 (时间序列步数)
|
|
||||||
STEP_SIZE = 20 # 滑动窗口的步长
|
|
||||||
BATCH_SIZE = 256
|
|
||||||
|
|
||||||
# Features Configuration
|
|
||||||
INPUT_SENSOR = 'WSMS00012'
|
|
||||||
OUTPUT_SENSORS = ['WSMS00007', 'WSMS00008', 'WSMS00009', 'WSMS00010', 'WSMS00011']
|
|
||||||
INPUT_AXIS = 'value1' # 底部传感器输入轴(课程要求:X轴)
|
|
||||||
OUTPUT_AXIS = 'value3' # 目标传感器输出轴
|
|
||||||
|
|
||||||
# Model Configuration
|
TaskMode = Literal["forward", "inverse"]
|
||||||
CHANNELS = [64, 64, 128, 128, 256, 256] # TCN 各层通道数
|
IncompleteFilePolicy = Literal["skip", "raise"]
|
||||||
KERNEL_SIZE = 5
|
|
||||||
DROPOUT = 0.2
|
|
||||||
|
|
||||||
# Training Configuration
|
|
||||||
LEARNING_RATE = 1e-4
|
|
||||||
EPOCHS = 50
|
|
||||||
WEIGHT_DECAY = 1e-3
|
|
||||||
ENABLE_EARLY_STOP = False
|
|
||||||
EARLY_STOP_PATIENCE = 15
|
|
||||||
SPECTRAL_LOSS_WEIGHT = 0.6
|
|
||||||
CORR_LOSS_WEIGHT = 0.3
|
|
||||||
CORR_LOSS_EPS = 1e-8
|
|
||||||
|
|
||||||
# Device
|
@dataclass
|
||||||
import torch
|
class DataConfig:
|
||||||
DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
|
project_root: Path = field(default_factory=lambda: Path(__file__).resolve().parents[1])
|
||||||
|
scenario: str = "Non_TMD"
|
||||||
|
task: TaskMode = "forward"
|
||||||
|
|
||||||
|
train_split_name: str = "train"
|
||||||
|
val_split_name: str = "val"
|
||||||
|
test_split_name: str = "test"
|
||||||
|
csv_pattern: str = "*.csv"
|
||||||
|
|
||||||
|
code_column: str = "code"
|
||||||
|
time_column: str = "time"
|
||||||
|
base_sensor_code: str = "WSMS00012"
|
||||||
|
base_axis: str = "value1"
|
||||||
|
response_sensor_code: str = "WSMS00007"
|
||||||
|
response_axis: str = "value3"
|
||||||
|
|
||||||
|
window_size: int = 512
|
||||||
|
window_stride: int = 64
|
||||||
|
min_sequence_length: int = 512
|
||||||
|
|
||||||
|
batch_size: int = 32
|
||||||
|
num_workers: int = 0
|
||||||
|
pin_memory: bool = True
|
||||||
|
shuffle_train: bool = True
|
||||||
|
drop_last_train: bool = False
|
||||||
|
|
||||||
|
interpolation_method: str = "linear"
|
||||||
|
incomplete_file_policy: IncompleteFilePolicy = "skip"
|
||||||
|
normalization_eps: float = 1e-6
|
||||||
|
|
||||||
|
downloads_dir: Path = field(init=False)
|
||||||
|
scenario_dir: Path = field(init=False)
|
||||||
|
train_dir: Path = field(init=False)
|
||||||
|
val_dir: Path = field(init=False)
|
||||||
|
test_dir: Path = field(init=False)
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
self.project_root = Path(self.project_root).resolve()
|
||||||
|
self.downloads_dir = self.project_root / "downloads"
|
||||||
|
self.scenario_dir = self.downloads_dir / self.scenario
|
||||||
|
self.train_dir = self.scenario_dir / self.train_split_name
|
||||||
|
self.val_dir = self.scenario_dir / self.val_split_name
|
||||||
|
self.test_dir = self.scenario_dir / self.test_split_name
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ModelConfig:
|
||||||
|
input_channels: int = 1
|
||||||
|
output_channels: int = 1
|
||||||
|
tcn_channels: tuple[int, ...] = (32, 32, 64, 64, 128)
|
||||||
|
kernel_size: int = 5
|
||||||
|
dropout: float = 0.25
|
||||||
|
use_causal_conv: bool = True
|
||||||
|
dilation_base: int = 2
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LossConfig:
|
||||||
|
time_loss: str = "huber"
|
||||||
|
huber_delta: float = 1.0
|
||||||
|
fft_loss_weight: float = 0.005
|
||||||
|
fft_loss_min_hz: float = 0.1
|
||||||
|
fft_loss_max_hz: float = 5.0
|
||||||
|
forward_time_loss_weight: float = 0.1
|
||||||
|
forward_rms_loss_weight: float = 8.0
|
||||||
|
forward_scale_loss_weight: float = 8.0
|
||||||
|
forward_underestimate_loss_weight: float = 16.0
|
||||||
|
forward_envelope_loss_weight: float = 8.0
|
||||||
|
envelope_kernel_size: int = 33
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TrainConfig:
|
||||||
|
epochs: int = 100
|
||||||
|
learning_rate: float = 1e-3
|
||||||
|
weight_decay: float = 1e-4
|
||||||
|
seed: int = 42
|
||||||
|
device: str = "cuda"
|
||||||
|
grad_clip_norm: float = 1.0
|
||||||
|
use_amp: bool = True
|
||||||
|
lr_scheduler_patience: int = 5
|
||||||
|
lr_scheduler_factor: float = 0.5
|
||||||
|
min_learning_rate: float = 1e-6
|
||||||
|
early_stop_patience: int = 8
|
||||||
|
checkpoint_dir: str = "checkpoints"
|
||||||
|
best_model_name: str = "best_tcn_model.pt"
|
||||||
|
history_name: str = "training_history.csv"
|
||||||
|
dominant_freq_min_hz: float = 0.1
|
||||||
|
dominant_freq_max_hz: float = 5.0
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ExperimentConfig:
|
||||||
|
data: DataConfig = field(default_factory=DataConfig)
|
||||||
|
model: ModelConfig = field(default_factory=ModelConfig)
|
||||||
|
loss: LossConfig = field(default_factory=LossConfig)
|
||||||
|
train: TrainConfig = field(default_factory=TrainConfig)
|
||||||
|
|
||||||
|
|
||||||
|
def make_forward_config() -> ExperimentConfig:
|
||||||
|
config = ExperimentConfig()
|
||||||
|
config.data.task = "forward"
|
||||||
|
config.model.input_channels = 1
|
||||||
|
config.model.output_channels = 1
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
|
def make_inverse_config() -> ExperimentConfig:
|
||||||
|
config = ExperimentConfig()
|
||||||
|
config.data.task = "inverse"
|
||||||
|
config.model.input_channels = 1
|
||||||
|
config.model.output_channels = 1
|
||||||
|
return config
|
||||||
|
|||||||
517
src/dataset.py
517
src/dataset.py
@@ -1,188 +1,393 @@
|
|||||||
import os
|
from __future__ import annotations
|
||||||
import glob
|
|
||||||
import pandas as pd
|
from dataclasses import dataclass, field
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
import torch
|
import torch
|
||||||
from torch.utils.data import Dataset, DataLoader
|
from torch.utils.data import DataLoader, Dataset
|
||||||
from sklearn.preprocessing import StandardScaler
|
|
||||||
from config import *
|
|
||||||
|
|
||||||
class MultiOutputStandardizer:
|
try:
|
||||||
"""Per-channel standardization that ignores missing labels via masks."""
|
from .config import DataConfig, ExperimentConfig
|
||||||
def __init__(self, n_outputs):
|
except ImportError:
|
||||||
self.n_outputs = n_outputs
|
from config import DataConfig, ExperimentConfig
|
||||||
self.mean_ = np.zeros(n_outputs, dtype=np.float32)
|
|
||||||
self.scale_ = np.ones(n_outputs, dtype=np.float32)
|
|
||||||
self.fitted = False
|
|
||||||
|
|
||||||
def fit(self, y_sequences, mask_sequences):
|
|
||||||
means = []
|
|
||||||
scales = []
|
|
||||||
for c in range(self.n_outputs):
|
|
||||||
valid_values = []
|
|
||||||
for y_seq, m_seq in zip(y_sequences, mask_sequences):
|
|
||||||
valid = m_seq[:, c] > 0.5
|
|
||||||
if np.any(valid):
|
|
||||||
valid_values.append(y_seq[valid, c])
|
|
||||||
if len(valid_values) == 0:
|
|
||||||
means.append(0.0)
|
|
||||||
scales.append(1.0)
|
|
||||||
continue
|
|
||||||
vals = np.concatenate(valid_values, axis=0)
|
|
||||||
mean = float(np.mean(vals))
|
|
||||||
std = float(np.std(vals))
|
|
||||||
if std < 1e-6:
|
|
||||||
std = 1.0
|
|
||||||
means.append(mean)
|
|
||||||
scales.append(std)
|
|
||||||
|
|
||||||
self.mean_ = np.asarray(means, dtype=np.float32)
|
|
||||||
self.scale_ = np.asarray(scales, dtype=np.float32)
|
|
||||||
self.fitted = True
|
|
||||||
|
|
||||||
def transform(self, y):
|
|
||||||
if not self.fitted:
|
|
||||||
raise RuntimeError("MultiOutputStandardizer must be fitted before transform.")
|
|
||||||
return (y - self.mean_) / self.scale_
|
|
||||||
|
|
||||||
def inverse_transform(self, y):
|
|
||||||
if not self.fitted:
|
|
||||||
raise RuntimeError("MultiOutputStandardizer must be fitted before inverse_transform.")
|
|
||||||
return y * self.scale_ + self.mean_
|
|
||||||
|
|
||||||
|
|
||||||
class BuildingDataset(Dataset):
|
@dataclass
|
||||||
def __init__(self, file_paths, seq_len, step_size, scaler_X=None, scaler_Y=None, fit_scaler=False):
|
class NormalizationStats:
|
||||||
self.seq_len = seq_len
|
x_mean: torch.Tensor
|
||||||
self.X_data = []
|
x_std: torch.Tensor
|
||||||
self.Y_data = []
|
y_mean: torch.Tensor
|
||||||
self.M_data = []
|
y_std: torch.Tensor
|
||||||
self.scaler_X = scaler_X if scaler_X is not None else StandardScaler()
|
|
||||||
self.scaler_Y = scaler_Y if scaler_Y is not None else MultiOutputStandardizer(len(OUTPUT_SENSORS))
|
|
||||||
|
|
||||||
raw_X = []
|
def normalize_x(self, tensor: torch.Tensor) -> torch.Tensor:
|
||||||
raw_Y = []
|
return (tensor - self.x_mean) / self.x_std
|
||||||
raw_M = []
|
|
||||||
|
|
||||||
for f in file_paths:
|
def denormalize_x(self, tensor: torch.Tensor) -> torch.Tensor:
|
||||||
# 读取数据
|
return tensor * self.x_std + self.x_mean
|
||||||
df = pd.read_csv(f)
|
|
||||||
# 使用长表格式: code, type, time, value1, value2, value3
|
|
||||||
|
|
||||||
# 提取 012 的输入轴作为基准
|
def normalize_y(self, tensor: torch.Tensor) -> torch.Tensor:
|
||||||
df_in = df[df['code'] == INPUT_SENSOR][['time', INPUT_AXIS]].rename(columns={INPUT_AXIS: 'input_signal'})
|
return (tensor - self.y_mean) / self.y_std
|
||||||
if len(df_in) == 0:
|
|
||||||
# 若无输入传感器(如自由衰减数据),则补零
|
|
||||||
df_in = pd.DataFrame({'time': df['time'].unique()})
|
|
||||||
df_in['input_signal'] = 0.0
|
|
||||||
|
|
||||||
# 提取 007~011 的 Z 轴并按时间戳逐步合并 (使用 left join 确保以 df_in 的时间为基准)
|
def denormalize_y(self, tensor: torch.Tensor) -> torch.Tensor:
|
||||||
df_merged = df_in
|
return tensor * self.y_std + self.y_mean
|
||||||
for sens in OUTPUT_SENSORS:
|
|
||||||
df_out_sens = df[df['code'] == sens][['time', OUTPUT_AXIS]].rename(columns={OUTPUT_AXIS: f'out_{sens}'})
|
|
||||||
df_merged = pd.merge(df_merged, df_out_sens, on='time', how='left')
|
|
||||||
|
|
||||||
df_merged = df_merged.sort_values('time').reset_index(drop=True)
|
|
||||||
|
|
||||||
if len(df_merged) == 0:
|
@dataclass
|
||||||
print(f"Warning: Skipping file {f} due to no overlapping timestamps across required sensors.")
|
class TensorStandardScaler:
|
||||||
|
mean: torch.Tensor
|
||||||
|
std: torch.Tensor
|
||||||
|
|
||||||
|
def transform(self, array: torch.Tensor | np.ndarray) -> torch.Tensor | np.ndarray:
|
||||||
|
return self._apply(array, inverse=False)
|
||||||
|
|
||||||
|
def inverse_transform(self, array: torch.Tensor | np.ndarray) -> torch.Tensor | np.ndarray:
|
||||||
|
return self._apply(array, inverse=True)
|
||||||
|
|
||||||
|
def _apply(self, array: torch.Tensor | np.ndarray, inverse: bool) -> torch.Tensor | np.ndarray:
|
||||||
|
if isinstance(array, torch.Tensor):
|
||||||
|
mean = self.mean.to(array.device, dtype=array.dtype)
|
||||||
|
std = self.std.to(array.device, dtype=array.dtype)
|
||||||
|
return array * std + mean if inverse else (array - mean) / std
|
||||||
|
|
||||||
|
np_array = np.asarray(array, dtype=np.float32)
|
||||||
|
mean_np = self.mean.detach().cpu().numpy().astype(np.float32)
|
||||||
|
std_np = self.std.detach().cpu().numpy().astype(np.float32)
|
||||||
|
return np_array * std_np + mean_np if inverse else (np_array - mean_np) / std_np
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class SequenceRecord:
|
||||||
|
file_path: Path
|
||||||
|
split: str
|
||||||
|
time: torch.Tensor
|
||||||
|
x_raw: torch.Tensor
|
||||||
|
y_raw: torch.Tensor
|
||||||
|
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)
|
||||||
|
short_files: list[tuple[str, int]] = field(default_factory=list)
|
||||||
|
|
||||||
|
def to_lines(self) -> list[str]:
|
||||||
|
lines = [
|
||||||
|
f"[{self.split}] loaded={len(self.loaded_files)} "
|
||||||
|
f"skipped={len(self.skipped_files)} short={len(self.short_files)}"
|
||||||
|
]
|
||||||
|
if self.skipped_files:
|
||||||
|
for file_name, reason in self.skipped_files:
|
||||||
|
lines.append(f" - skipped {file_name}: {reason}")
|
||||||
|
if self.short_files:
|
||||||
|
for file_name, length in self.short_files:
|
||||||
|
lines.append(f" - short {file_name}: length={length}")
|
||||||
|
if self.interpolated_files:
|
||||||
|
for file_name, count in self.interpolated_files.items():
|
||||||
|
if count > 0:
|
||||||
|
lines.append(f" - interpolated {file_name}: missing_points={count}")
|
||||||
|
return lines
|
||||||
|
|
||||||
|
|
||||||
|
class WindowedTimeSeriesDataset(Dataset):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
records: list[SequenceRecord],
|
||||||
|
normalization: NormalizationStats,
|
||||||
|
config: DataConfig,
|
||||||
|
split: str,
|
||||||
|
report: SplitLoadReport | None = None,
|
||||||
|
) -> None:
|
||||||
|
self.config = config
|
||||||
|
self.split = split
|
||||||
|
self.normalization = normalization
|
||||||
|
self.report = report or SplitLoadReport(split=split)
|
||||||
|
self.sequence_store: list[dict[str, Any]] = []
|
||||||
|
self.window_index: list[tuple[int, int]] = []
|
||||||
|
|
||||||
|
min_length = max(config.min_sequence_length, config.window_size)
|
||||||
|
|
||||||
|
for record in records:
|
||||||
|
sequence_length = int(record.time.shape[0])
|
||||||
|
if sequence_length < min_length:
|
||||||
|
self.report.short_files.append((record.file_path.name, sequence_length))
|
||||||
continue
|
continue
|
||||||
|
|
||||||
x_seq = df_merged['input_signal'].values.reshape(-1, 1).astype(np.float32)
|
x_norm = normalization.normalize_x(record.x_raw)
|
||||||
|
y_norm = normalization.normalize_y(record.y_raw)
|
||||||
|
|
||||||
# 提取所有 target 传感器列与可用性掩码
|
record_index = len(self.sequence_store)
|
||||||
out_cols = [f'out_{sens}' for sens in OUTPUT_SENSORS]
|
self.sequence_store.append(
|
||||||
y_seq = np.zeros((len(df_merged), len(OUTPUT_SENSORS)), dtype=np.float32)
|
{
|
||||||
m_seq = np.zeros((len(df_merged), len(OUTPUT_SENSORS)), dtype=np.float32)
|
"file_name": record.file_path.name,
|
||||||
for c, col in enumerate(out_cols):
|
"file_path": str(record.file_path),
|
||||||
series = df_merged[col]
|
"time": record.time,
|
||||||
observed = ~series.isna()
|
"x": x_norm,
|
||||||
m_seq[:, c] = observed.astype(np.float32)
|
"y": y_norm,
|
||||||
if observed.any():
|
}
|
||||||
filled = series.interpolate(method='linear').bfill().ffill()
|
)
|
||||||
y_seq[:, c] = filled.fillna(0.0).values.astype(np.float32)
|
|
||||||
else:
|
|
||||||
y_seq[:, c] = 0.0
|
|
||||||
|
|
||||||
raw_X.append(x_seq)
|
last_start = sequence_length - config.window_size
|
||||||
raw_Y.append(y_seq)
|
for start in range(0, last_start + 1, config.window_stride):
|
||||||
raw_M.append(m_seq)
|
self.window_index.append((record_index, start))
|
||||||
|
|
||||||
if len(raw_X) == 0:
|
if not self.window_index:
|
||||||
raise ValueError("未能从文件中构造出有效序列,请检查数据路径与传感器编码配置。")
|
raise RuntimeError(
|
||||||
|
f"No valid windows were created for split='{split}'. "
|
||||||
|
"Check file completeness, sequence length, and window configuration."
|
||||||
|
)
|
||||||
|
|
||||||
# 拼接所有文件数据进行 fit
|
def __len__(self) -> int:
|
||||||
X_all = np.vstack(raw_X)
|
return len(self.window_index)
|
||||||
|
|
||||||
if fit_scaler:
|
def __getitem__(self, index: int) -> dict[str, Any]:
|
||||||
self.scaler_X.fit(X_all)
|
record_index, start = self.window_index[index]
|
||||||
self.scaler_Y.fit(raw_Y, raw_M)
|
record = self.sequence_store[record_index]
|
||||||
|
end = start + self.config.window_size
|
||||||
|
|
||||||
# 切分窗口
|
x = record["x"][start:end]
|
||||||
for x_seq, y_seq, m_seq in zip(raw_X, raw_Y, raw_M):
|
time = record["time"][start:end]
|
||||||
x_seq_scaled = self.scaler_X.transform(x_seq)
|
y = record["y"][start:end]
|
||||||
y_seq_scaled = self.scaler_Y.transform(y_seq)
|
|
||||||
y_seq_scaled = np.where(m_seq > 0.5, y_seq_scaled, 0.0).astype(np.float32)
|
|
||||||
|
|
||||||
for i in range(0, len(x_seq_scaled) - seq_len + 1, step_size):
|
return {
|
||||||
x_win = x_seq_scaled[i:i+seq_len]
|
"x": x,
|
||||||
y_win = y_seq_scaled[i:i+seq_len]
|
"y": y,
|
||||||
m_win = m_seq[i:i+seq_len]
|
"time": time,
|
||||||
if np.sum(m_win) <= 0:
|
"file_name": record["file_name"],
|
||||||
continue
|
"file_path": record["file_path"],
|
||||||
self.X_data.append(x_win)
|
"window_start": torch.tensor(start, dtype=torch.long),
|
||||||
self.Y_data.append(y_win)
|
"window_end": torch.tensor(end, dtype=torch.long),
|
||||||
self.M_data.append(m_win)
|
}
|
||||||
|
|
||||||
self.X_data = np.array(self.X_data)
|
def denormalize_x(self, tensor: torch.Tensor) -> torch.Tensor:
|
||||||
self.Y_data = np.array(self.Y_data)
|
return self.normalization.denormalize_x(tensor)
|
||||||
self.M_data = np.array(self.M_data)
|
|
||||||
|
|
||||||
def __len__(self):
|
def denormalize_y(self, tensor: torch.Tensor) -> torch.Tensor:
|
||||||
return len(self.X_data)
|
return self.normalization.denormalize_y(tensor)
|
||||||
|
|
||||||
def __getitem__(self, idx):
|
|
||||||
return (
|
def list_split_files(config: DataConfig, split: str) -> list[Path]:
|
||||||
torch.tensor(self.X_data[idx], dtype=torch.float32),
|
split_dir = getattr(config, f"{split}_dir")
|
||||||
torch.tensor(self.Y_data[idx], dtype=torch.float32),
|
return sorted(split_dir.glob(config.csv_pattern))
|
||||||
torch.tensor(self.M_data[idx], dtype=torch.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 _handle_incomplete_file(
|
||||||
|
file_path: Path,
|
||||||
|
split: str,
|
||||||
|
reason: str,
|
||||||
|
config: DataConfig,
|
||||||
|
report: SplitLoadReport,
|
||||||
|
) -> None:
|
||||||
|
if config.incomplete_file_policy == "raise":
|
||||||
|
raise ValueError(f"{file_path.name}: {reason}")
|
||||||
|
report.skipped_files.append((file_path.name, reason))
|
||||||
|
|
||||||
|
|
||||||
|
def align_file_to_base_timeline(
|
||||||
|
file_path: Path,
|
||||||
|
config: DataConfig,
|
||||||
|
split: str,
|
||||||
|
report: SplitLoadReport,
|
||||||
|
) -> SequenceRecord | None:
|
||||||
|
df = pd.read_csv(file_path)
|
||||||
|
|
||||||
|
base_df = _value_frame(df, config, config.base_sensor_code, config.base_axis)
|
||||||
|
if base_df.empty:
|
||||||
|
_handle_incomplete_file(
|
||||||
|
file_path=file_path,
|
||||||
|
split=split,
|
||||||
|
reason=f"missing base sensor {config.base_sensor_code}",
|
||||||
|
config=config,
|
||||||
|
report=report,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
response_df = _value_frame(df, config, config.response_sensor_code, config.response_axis)
|
||||||
|
if response_df.empty:
|
||||||
|
_handle_incomplete_file(
|
||||||
|
file_path=file_path,
|
||||||
|
split=split,
|
||||||
|
reason=f"missing response sensor {config.response_sensor_code}",
|
||||||
|
config=config,
|
||||||
|
report=report,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
aligned = base_df.rename(columns={config.base_axis: "base_signal"})
|
||||||
|
response_df = response_df.rename(columns={config.response_axis: "response_signal"})
|
||||||
|
aligned = aligned.merge(response_df, 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",
|
||||||
|
)
|
||||||
|
aligned[["response_signal"]] = aligned[["response_signal"]].ffill().bfill()
|
||||||
|
|
||||||
|
if aligned["response_signal"].isna().any():
|
||||||
|
_handle_incomplete_file(
|
||||||
|
file_path=file_path,
|
||||||
|
split=split,
|
||||||
|
reason="remaining NaN values after alignment/interpolation",
|
||||||
|
config=config,
|
||||||
|
report=report,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
if interpolation_count > 0:
|
||||||
|
report.interpolated_files[file_path.name] = interpolation_count
|
||||||
|
|
||||||
|
time_tensor = torch.tensor(aligned[config.time_column].to_numpy(), dtype=torch.float64)
|
||||||
|
base_tensor = torch.tensor(aligned[["base_signal"]].to_numpy(), dtype=torch.float32)
|
||||||
|
response_tensor = torch.tensor(aligned[["response_signal"]].to_numpy(), dtype=torch.float32)
|
||||||
|
|
||||||
|
if config.task == "forward":
|
||||||
|
x_raw = base_tensor
|
||||||
|
y_raw = response_tensor
|
||||||
|
elif config.task == "inverse":
|
||||||
|
x_raw = response_tensor
|
||||||
|
y_raw = base_tensor
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported task: {config.task}")
|
||||||
|
|
||||||
|
report.loaded_files.append(file_path.name)
|
||||||
|
return SequenceRecord(
|
||||||
|
file_path=file_path,
|
||||||
|
split=split,
|
||||||
|
time=time_tensor,
|
||||||
|
x_raw=x_raw,
|
||||||
|
y_raw=y_raw,
|
||||||
|
interpolation_count=interpolation_count,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def load_split_records(config: DataConfig, split: str) -> tuple[list[SequenceRecord], SplitLoadReport]:
|
||||||
|
report = SplitLoadReport(split=split)
|
||||||
|
records: list[SequenceRecord] = []
|
||||||
|
|
||||||
|
for file_path in list_split_files(config, split):
|
||||||
|
record = align_file_to_base_timeline(file_path=file_path, config=config, split=split, report=report)
|
||||||
|
if record is not None:
|
||||||
|
records.append(record)
|
||||||
|
|
||||||
|
if not records:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"No usable files found in split='{split}' under '{getattr(config, f'{split}_dir')}'."
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_dataloaders(condition='Non_TMD', include_free_vib=False):
|
return records, report
|
||||||
"""
|
|
||||||
condition: 'Non_TMD' 或者是 'TMD'
|
|
||||||
include_free_vib: 是否在训练集中加入自由振动与自由衰减数据
|
|
||||||
"""
|
|
||||||
base_dir = os.path.join(DATA_DIR, condition)
|
|
||||||
|
|
||||||
# 手动区分的子目录
|
|
||||||
train_files = glob.glob(os.path.join(base_dir, 'train', '*.csv'))
|
|
||||||
if include_free_vib:
|
|
||||||
train_files += glob.glob(os.path.join(base_dir, 'free_vib', '*.csv'))
|
|
||||||
|
|
||||||
val_files = glob.glob(os.path.join(base_dir, 'val', '*.csv'))
|
def _safe_feature_std(tensor: torch.Tensor, eps: float) -> torch.Tensor:
|
||||||
test_files = glob.glob(os.path.join(base_dir, 'test', '*.csv'))
|
std = tensor.std(dim=0, unbiased=False)
|
||||||
|
eps_tensor = torch.full_like(std, eps)
|
||||||
|
return torch.maximum(std, eps_tensor)
|
||||||
|
|
||||||
print(f"[{condition}] Train files: {len(train_files)}, Val files: {len(val_files)}, Test files: {len(test_files)}")
|
|
||||||
|
|
||||||
if len(train_files) == 0:
|
def fit_normalization_stats(records: list[SequenceRecord], config: DataConfig) -> NormalizationStats:
|
||||||
raise ValueError(f"错误: 在 {base_dir}/train 目录下未找到训练文件!请检查路径是否正确。")
|
x_all = torch.cat([record.x_raw for record in records], dim=0)
|
||||||
if len(val_files) == 0:
|
y_all = torch.cat([record.y_raw for record in records], dim=0)
|
||||||
raise ValueError(f"错误: 在 {base_dir}/val 目录下未找到验证文件!请检查路径是否正确。")
|
x_mean = x_all.mean(dim=0)
|
||||||
|
x_std = _safe_feature_std(x_all, config.normalization_eps)
|
||||||
|
y_mean = y_all.mean(dim=0)
|
||||||
|
y_std = _safe_feature_std(y_all, config.normalization_eps)
|
||||||
|
|
||||||
train_dataset = BuildingDataset(train_files, SEQ_LEN, STEP_SIZE, fit_scaler=True)
|
return NormalizationStats(
|
||||||
val_dataset = BuildingDataset(val_files, SEQ_LEN, STEP_SIZE,
|
x_mean=x_mean,
|
||||||
scaler_X=train_dataset.scaler_X,
|
x_std=x_std,
|
||||||
scaler_Y=train_dataset.scaler_Y, fit_scaler=False)
|
y_mean=y_mean,
|
||||||
# 若某条件(如 TMD)下没有 test 数据,可以处理一下防止报错
|
y_std=y_std,
|
||||||
test_loader = None
|
)
|
||||||
if len(test_files) > 0:
|
|
||||||
test_dataset = BuildingDataset(test_files, SEQ_LEN, STEP_SIZE,
|
|
||||||
scaler_X=train_dataset.scaler_X,
|
|
||||||
scaler_Y=train_dataset.scaler_Y, fit_scaler=False)
|
|
||||||
test_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False)
|
|
||||||
|
|
||||||
train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)
|
|
||||||
val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False)
|
|
||||||
|
|
||||||
return train_loader, val_loader, test_loader, train_dataset.scaler_X, train_dataset.scaler_Y
|
def build_datasets(
|
||||||
|
config: ExperimentConfig | DataConfig,
|
||||||
|
) -> tuple[dict[str, WindowedTimeSeriesDataset], NormalizationStats, 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_stats(train_records, data_config)
|
||||||
|
|
||||||
|
val_records, val_report = load_split_records(data_config, "val")
|
||||||
|
test_records, test_report = load_split_records(data_config, "test")
|
||||||
|
|
||||||
|
datasets = {
|
||||||
|
"train": WindowedTimeSeriesDataset(train_records, normalization, data_config, "train", train_report),
|
||||||
|
"val": WindowedTimeSeriesDataset(val_records, normalization, data_config, "val", val_report),
|
||||||
|
"test": WindowedTimeSeriesDataset(test_records, normalization, data_config, "test", test_report),
|
||||||
|
}
|
||||||
|
reports = {"train": train_report, "val": val_report, "test": test_report}
|
||||||
|
return datasets, normalization, reports
|
||||||
|
|
||||||
|
|
||||||
|
def build_dataloaders(
|
||||||
|
config: ExperimentConfig | DataConfig,
|
||||||
|
) -> tuple[dict[str, DataLoader], dict[str, WindowedTimeSeriesDataset], dict[str, SplitLoadReport]]:
|
||||||
|
data_config = config.data if isinstance(config, ExperimentConfig) else config
|
||||||
|
datasets, _, reports = build_datasets(config)
|
||||||
|
|
||||||
|
loaders = {
|
||||||
|
"train": DataLoader(
|
||||||
|
datasets["train"],
|
||||||
|
batch_size=data_config.batch_size,
|
||||||
|
shuffle=data_config.shuffle_train,
|
||||||
|
num_workers=data_config.num_workers,
|
||||||
|
pin_memory=data_config.pin_memory,
|
||||||
|
drop_last=data_config.drop_last_train,
|
||||||
|
),
|
||||||
|
"val": DataLoader(
|
||||||
|
datasets["val"],
|
||||||
|
batch_size=data_config.batch_size,
|
||||||
|
shuffle=False,
|
||||||
|
num_workers=data_config.num_workers,
|
||||||
|
pin_memory=data_config.pin_memory,
|
||||||
|
drop_last=False,
|
||||||
|
),
|
||||||
|
"test": DataLoader(
|
||||||
|
datasets["test"],
|
||||||
|
batch_size=data_config.batch_size,
|
||||||
|
shuffle=False,
|
||||||
|
num_workers=data_config.num_workers,
|
||||||
|
pin_memory=data_config.pin_memory,
|
||||||
|
drop_last=False,
|
||||||
|
),
|
||||||
|
}
|
||||||
|
return loaders, datasets, reports
|
||||||
|
|
||||||
|
|
||||||
|
def get_dataloaders(
|
||||||
|
config: ExperimentConfig | DataConfig,
|
||||||
|
) -> tuple[
|
||||||
|
dict[str, DataLoader],
|
||||||
|
dict[str, WindowedTimeSeriesDataset],
|
||||||
|
dict[str, SplitLoadReport],
|
||||||
|
TensorStandardScaler,
|
||||||
|
TensorStandardScaler,
|
||||||
|
]:
|
||||||
|
loaders, datasets, reports = build_dataloaders(config)
|
||||||
|
normalization = datasets["train"].normalization
|
||||||
|
scaler_x = TensorStandardScaler(mean=normalization.x_mean, std=normalization.x_std)
|
||||||
|
scaler_y = TensorStandardScaler(mean=normalization.y_mean, std=normalization.y_std)
|
||||||
|
return loaders, datasets, reports, scaler_x, scaler_y
|
||||||
|
|
||||||
|
|
||||||
|
def report_to_text(reports: dict[str, SplitLoadReport]) -> str:
|
||||||
|
lines: list[str] = []
|
||||||
|
for split in ("train", "val", "test"):
|
||||||
|
report = reports[split]
|
||||||
|
lines.extend(report.to_lines())
|
||||||
|
return "\n".join(lines)
|
||||||
|
|||||||
298
src/evaluate.py
298
src/evaluate.py
@@ -1,50 +1,274 @@
|
|||||||
import torch
|
from __future__ import annotations
|
||||||
import numpy as np
|
|
||||||
|
import argparse
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
import matplotlib.pyplot as plt
|
import matplotlib.pyplot as plt
|
||||||
from config import *
|
import numpy as np
|
||||||
from dataset import get_dataloaders
|
import torch
|
||||||
from model import BuildingTCN
|
|
||||||
|
|
||||||
def evaluate_model():
|
try:
|
||||||
_, _, test_loader, scaler_X, scaler_Y = get_dataloaders()
|
from .config import ExperimentConfig, make_forward_config, make_inverse_config
|
||||||
if test_loader is None:
|
from .dataset import get_dataloaders, report_to_text
|
||||||
raise ValueError("当前数据配置下没有可用的测试集。")
|
from .model import build_model
|
||||||
|
except ImportError:
|
||||||
|
from config import ExperimentConfig, make_forward_config, make_inverse_config
|
||||||
|
from dataset import get_dataloaders, report_to_text
|
||||||
|
from model import build_model
|
||||||
|
|
||||||
model = BuildingTCN(input_size=1, output_size=5, num_channels=CHANNELS,
|
|
||||||
kernel_size=KERNEL_SIZE, dropout=DROPOUT).to(DEVICE)
|
def calculate_rms(signal: np.ndarray) -> float:
|
||||||
model.load_state_dict(torch.load('best_model.pth', map_location=DEVICE))
|
signal = np.asarray(signal, dtype=np.float64).reshape(-1)
|
||||||
|
if signal.size == 0:
|
||||||
|
raise ValueError("signal must not be empty")
|
||||||
|
return float(np.sqrt(np.mean(np.square(signal))))
|
||||||
|
|
||||||
|
|
||||||
|
def get_dominant_frequency(signal: np.ndarray, sampling_rate: float) -> float:
|
||||||
|
signal = np.asarray(signal, dtype=np.float64).reshape(-1)
|
||||||
|
if signal.size == 0:
|
||||||
|
raise ValueError("signal must not be empty")
|
||||||
|
fft_values = np.fft.rfft(signal)
|
||||||
|
frequencies = np.fft.rfftfreq(signal.size, d=1.0 / sampling_rate)
|
||||||
|
magnitudes = np.abs(fft_values)
|
||||||
|
if magnitudes.size > 0:
|
||||||
|
magnitudes[0] = 0.0
|
||||||
|
dominant_index = int(np.argmax(magnitudes))
|
||||||
|
return float(frequencies[dominant_index])
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser(description="Evaluate trained TCN with time-domain and frequency-domain plots.")
|
||||||
|
parser.add_argument("--task", choices=("forward", "inverse"), default="forward")
|
||||||
|
parser.add_argument("--split", choices=("val", "test"), default="test")
|
||||||
|
parser.add_argument("--batch-index", type=int, default=0)
|
||||||
|
parser.add_argument("--sample-index", type=int, default=0)
|
||||||
|
parser.add_argument("--sampling-rate", type=float, default=None)
|
||||||
|
parser.add_argument("--device", type=str, default=None)
|
||||||
|
parser.add_argument("--checkpoint", type=str, default=None)
|
||||||
|
parser.add_argument("--save-dir", type=str, default=None)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def make_config(task: str) -> ExperimentConfig:
|
||||||
|
if task == "forward":
|
||||||
|
return make_forward_config()
|
||||||
|
if task == "inverse":
|
||||||
|
return make_inverse_config()
|
||||||
|
raise ValueError(f"Unsupported task: {task}")
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_device(device_name: str | None, config: ExperimentConfig) -> torch.device:
|
||||||
|
requested = device_name or config.train.device
|
||||||
|
if requested.startswith("cuda") and not torch.cuda.is_available():
|
||||||
|
return torch.device("cpu")
|
||||||
|
return torch.device(requested)
|
||||||
|
|
||||||
|
|
||||||
|
def default_checkpoint_path(config: ExperimentConfig) -> Path:
|
||||||
|
return config.data.project_root / config.train.checkpoint_dir / config.data.task / config.train.best_model_name
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_checkpoint_path(config: ExperimentConfig, checkpoint_arg: str | None) -> Path:
|
||||||
|
if checkpoint_arg:
|
||||||
|
return Path(checkpoint_arg).resolve()
|
||||||
|
|
||||||
|
default_path = default_checkpoint_path(config)
|
||||||
|
if default_path.exists():
|
||||||
|
return default_path
|
||||||
|
|
||||||
|
legacy_path = config.data.project_root / "best_model.pth"
|
||||||
|
if legacy_path.exists():
|
||||||
|
return legacy_path
|
||||||
|
|
||||||
|
return default_path
|
||||||
|
|
||||||
|
|
||||||
|
def select_batch(dataloader: torch.utils.data.DataLoader, batch_index: int) -> dict[str, object]:
|
||||||
|
for current_index, batch in enumerate(dataloader):
|
||||||
|
if current_index == batch_index:
|
||||||
|
return batch
|
||||||
|
raise IndexError(f"batch_index={batch_index} is out of range")
|
||||||
|
|
||||||
|
|
||||||
|
def get_sample_1d(batch_tensor: torch.Tensor, sample_index: int) -> torch.Tensor:
|
||||||
|
sample = batch_tensor[sample_index]
|
||||||
|
if sample.ndim == 2 and sample.shape[-1] == 1:
|
||||||
|
sample = sample.squeeze(-1)
|
||||||
|
return sample
|
||||||
|
|
||||||
|
|
||||||
|
def compute_fft_curve(signal: np.ndarray, sampling_rate: float) -> tuple[np.ndarray, np.ndarray]:
|
||||||
|
signal = np.asarray(signal, dtype=np.float64).reshape(-1)
|
||||||
|
fft_values = np.fft.rfft(signal)
|
||||||
|
frequencies = np.fft.rfftfreq(signal.size, d=1.0 / sampling_rate)
|
||||||
|
magnitudes = np.abs(fft_values)
|
||||||
|
return frequencies, magnitudes
|
||||||
|
|
||||||
|
|
||||||
|
def estimate_sampling_rate_from_time(time_values: np.ndarray) -> float:
|
||||||
|
time_values = np.asarray(time_values, dtype=np.float64).reshape(-1)
|
||||||
|
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
|
||||||
|
median_dt = float(np.median(dt))
|
||||||
|
if median_dt <= 0.0:
|
||||||
|
return 100.0
|
||||||
|
return 1.0 / median_dt
|
||||||
|
|
||||||
|
|
||||||
|
def relative_percent_error(true_value: float, pred_value: float) -> float:
|
||||||
|
denominator = max(abs(true_value), 1e-12)
|
||||||
|
return abs(pred_value - true_value) / denominator * 100.0
|
||||||
|
|
||||||
|
|
||||||
|
def task_labels(config: ExperimentConfig) -> tuple[str, str]:
|
||||||
|
if config.data.task == "forward":
|
||||||
|
return (
|
||||||
|
f"{config.data.base_sensor_code}/{config.data.base_axis}",
|
||||||
|
f"{config.data.response_sensor_code}/{config.data.response_axis}",
|
||||||
|
)
|
||||||
|
return (
|
||||||
|
f"{config.data.response_sensor_code}/{config.data.response_axis}",
|
||||||
|
f"{config.data.base_sensor_code}/{config.data.base_axis}",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def evaluate_one_batch(
|
||||||
|
config: ExperimentConfig,
|
||||||
|
split: str,
|
||||||
|
batch_index: int,
|
||||||
|
sample_index: int,
|
||||||
|
sampling_rate: float,
|
||||||
|
checkpoint_path: Path,
|
||||||
|
device: torch.device,
|
||||||
|
save_dir: Path,
|
||||||
|
) -> Path:
|
||||||
|
loaders, datasets, reports, scaler_x, scaler_y = get_dataloaders(config)
|
||||||
|
print(report_to_text(reports))
|
||||||
|
|
||||||
|
dataset = datasets[split]
|
||||||
|
dataloader = loaders[split]
|
||||||
|
batch = select_batch(dataloader, batch_index)
|
||||||
|
|
||||||
|
x = batch["x"].to(device)
|
||||||
|
y = batch["y"].to(device)
|
||||||
|
|
||||||
|
model = build_model(config).to(device)
|
||||||
|
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||||
|
model.load_state_dict(checkpoint["model_state_dict"])
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|
||||||
all_preds = []
|
|
||||||
all_targets = []
|
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
for inputs, targets, masks in test_loader:
|
pred_norm = model(x)
|
||||||
inputs = inputs.to(DEVICE)
|
|
||||||
outputs = model(inputs)
|
|
||||||
|
|
||||||
all_preds.append(outputs.cpu().numpy())
|
scaler_output = scaler_y if config.data.task == "forward" else scaler_x
|
||||||
all_targets.append(targets.numpy())
|
pred_physical = scaler_output.inverse_transform(pred_norm).detach().cpu()
|
||||||
|
true_physical = scaler_output.inverse_transform(y).detach().cpu()
|
||||||
|
|
||||||
# Concatenate results
|
pred_signal = get_sample_1d(pred_physical, sample_index).numpy()
|
||||||
all_preds = np.concatenate(all_preds, axis=0)
|
true_signal = get_sample_1d(true_physical, sample_index).numpy()
|
||||||
all_targets = np.concatenate(all_targets, axis=0)
|
time_values = batch["time"][sample_index].detach().cpu().numpy()
|
||||||
|
time_steps = np.arange(true_signal.shape[0], dtype=np.float64)
|
||||||
|
effective_sampling_rate = (
|
||||||
|
sampling_rate if sampling_rate is not None else estimate_sampling_rate_from_time(time_values)
|
||||||
|
)
|
||||||
|
|
||||||
# 取一个 batch 的第一条序列进行可视化
|
true_rms = calculate_rms(true_signal)
|
||||||
sample_pred = scaler_Y.inverse_transform(all_preds[0])
|
pred_rms = calculate_rms(pred_signal)
|
||||||
sample_target = scaler_Y.inverse_transform(all_targets[0])
|
rms_error_percent = relative_percent_error(true_rms, pred_rms)
|
||||||
|
|
||||||
plt.figure(figsize=(15, 10))
|
true_freq = get_dominant_frequency(true_signal, sampling_rate=effective_sampling_rate)
|
||||||
for i, sens in enumerate(OUTPUT_SENSORS):
|
pred_freq = get_dominant_frequency(pred_signal, sampling_rate=effective_sampling_rate)
|
||||||
plt.subplot(5, 1, i+1)
|
|
||||||
plt.plot(sample_target[:, i], label='True', alpha=0.7)
|
true_fft_freqs, true_fft_mag = compute_fft_curve(true_signal, sampling_rate=effective_sampling_rate)
|
||||||
plt.plot(sample_pred[:, i], label='Pred', alpha=0.7, linestyle='--')
|
pred_fft_freqs, pred_fft_mag = compute_fft_curve(pred_signal, sampling_rate=effective_sampling_rate)
|
||||||
plt.title(f'Sensor {sens} Z-axis Response (Test: Earthquake)')
|
|
||||||
plt.legend()
|
freq_mask_true = (true_fft_freqs >= 0.0) & (true_fft_freqs <= 5.0)
|
||||||
|
freq_mask_pred = (pred_fft_freqs >= 0.0) & (pred_fft_freqs <= 5.0)
|
||||||
|
|
||||||
|
input_label, output_label = task_labels(config)
|
||||||
|
file_name = batch["file_name"][sample_index]
|
||||||
|
window_start = int(batch["window_start"][sample_index].item())
|
||||||
|
|
||||||
|
fig, axes = plt.subplots(2, 1, figsize=(14, 9))
|
||||||
|
fig.suptitle(f"{config.data.task.upper()} evaluation | {split} | {file_name} | window_start={window_start}")
|
||||||
|
|
||||||
|
axes[0].plot(time_steps, true_signal, linewidth=1.2, label="True")
|
||||||
|
axes[0].plot(time_steps, pred_signal, linewidth=1.2, label="Pred")
|
||||||
|
axes[0].set_xlabel("Time Step")
|
||||||
|
axes[0].set_ylabel("Acceleration")
|
||||||
|
axes[0].set_title(
|
||||||
|
f"{output_label} | True RMS: {true_rms:.4f}, Pred RMS: {pred_rms:.4f}, Error: {rms_error_percent:.2f}%"
|
||||||
|
)
|
||||||
|
axes[0].grid(True, alpha=0.3)
|
||||||
|
axes[0].legend(loc="upper right")
|
||||||
|
|
||||||
|
axes[1].plot(true_fft_freqs[freq_mask_true], true_fft_mag[freq_mask_true], linewidth=1.2, label="True FFT")
|
||||||
|
axes[1].plot(pred_fft_freqs[freq_mask_pred], pred_fft_mag[freq_mask_pred], linewidth=1.2, label="Pred FFT")
|
||||||
|
true_peak_mag = np.interp(true_freq, true_fft_freqs, true_fft_mag)
|
||||||
|
pred_peak_mag = np.interp(pred_freq, pred_fft_freqs, pred_fft_mag)
|
||||||
|
if 0.0 <= true_freq <= 5.0:
|
||||||
|
axes[1].scatter([true_freq], [true_peak_mag], s=50, marker="o", label="True Peak")
|
||||||
|
if 0.0 <= pred_freq <= 5.0:
|
||||||
|
axes[1].scatter([pred_freq], [pred_peak_mag], s=50, marker="x", label="Pred Peak")
|
||||||
|
axes[1].set_xlim(0.0, 5.0)
|
||||||
|
axes[1].set_xlabel("Frequency (Hz)")
|
||||||
|
axes[1].set_ylabel("Amplitude")
|
||||||
|
axes[1].set_title(f"True Freq: {true_freq:.4f} Hz, Pred Freq: {pred_freq:.4f} Hz")
|
||||||
|
axes[1].grid(True, alpha=0.3)
|
||||||
|
axes[1].legend(loc="upper right")
|
||||||
|
|
||||||
plt.tight_layout()
|
plt.tight_layout()
|
||||||
plt.savefig('test_results.png')
|
save_dir.mkdir(parents=True, exist_ok=True)
|
||||||
print("Evaluation done. Result saved to test_results.png")
|
figure_path = save_dir / f"evaluation_{config.data.task}_{split}_b{batch_index}_s{sample_index}.png"
|
||||||
|
plt.savefig(figure_path, dpi=180, bbox_inches="tight")
|
||||||
|
plt.show()
|
||||||
|
|
||||||
if __name__ == '__main__':
|
print(f"Input channel: {input_label}")
|
||||||
evaluate_model()
|
print(f"Output channel: {output_label}")
|
||||||
|
print(f"Checkpoint: {checkpoint_path}")
|
||||||
|
print(f"Sampling Rate Used: {effective_sampling_rate:.6f} Hz")
|
||||||
|
print(f"Figure saved to: {figure_path}")
|
||||||
|
print(f"True RMS: {true_rms:.6f}")
|
||||||
|
print(f"Pred RMS: {pred_rms:.6f}")
|
||||||
|
print(f"RMS Error (%): {rms_error_percent:.4f}")
|
||||||
|
print(f"True Dominant Frequency: {true_freq:.6f} Hz")
|
||||||
|
print(f"Pred Dominant Frequency: {pred_freq:.6f} Hz")
|
||||||
|
print(f"Dominant Frequency Error: {abs(pred_freq - true_freq):.6f} Hz")
|
||||||
|
|
||||||
|
return figure_path
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
args = parse_args()
|
||||||
|
config = make_config(args.task)
|
||||||
|
device = resolve_device(args.device, config)
|
||||||
|
checkpoint_path = resolve_checkpoint_path(config, args.checkpoint)
|
||||||
|
save_dir = (
|
||||||
|
Path(args.save_dir).resolve()
|
||||||
|
if args.save_dir
|
||||||
|
else config.data.project_root / "evaluation_outputs" / config.data.task
|
||||||
|
)
|
||||||
|
|
||||||
|
if not checkpoint_path.exists():
|
||||||
|
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
|
||||||
|
|
||||||
|
evaluate_one_batch(
|
||||||
|
config=config,
|
||||||
|
split=args.split,
|
||||||
|
batch_index=args.batch_index,
|
||||||
|
sample_index=args.sample_index,
|
||||||
|
sampling_rate=args.sampling_rate,
|
||||||
|
checkpoint_path=checkpoint_path,
|
||||||
|
device=device,
|
||||||
|
save_dir=save_dir,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
|
|||||||
177
src/model.py
177
src/model.py
@@ -1,76 +1,151 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch.nn.utils import weight_norm
|
from torch.nn.utils import weight_norm
|
||||||
|
|
||||||
|
try:
|
||||||
|
from .config import ExperimentConfig, ModelConfig
|
||||||
|
except ImportError:
|
||||||
|
from config import ExperimentConfig, ModelConfig
|
||||||
|
|
||||||
|
|
||||||
class Chomp1d(nn.Module):
|
class Chomp1d(nn.Module):
|
||||||
def __init__(self, chomp_size):
|
def __init__(self, chomp_size: int) -> None:
|
||||||
super(Chomp1d, self).__init__()
|
super().__init__()
|
||||||
self.chomp_size = chomp_size
|
self.chomp_size = chomp_size
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
if self.chomp_size == 0:
|
||||||
|
return x
|
||||||
return x[:, :, :-self.chomp_size].contiguous()
|
return x[:, :, :-self.chomp_size].contiguous()
|
||||||
|
|
||||||
|
|
||||||
|
class CausalConv1d(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int,
|
||||||
|
out_channels: int,
|
||||||
|
kernel_size: int,
|
||||||
|
dilation: int = 1,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
padding = (kernel_size - 1) * dilation
|
||||||
|
self.net = nn.Sequential(
|
||||||
|
weight_norm(
|
||||||
|
nn.Conv1d(
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
kernel_size=kernel_size,
|
||||||
|
padding=padding,
|
||||||
|
dilation=dilation,
|
||||||
|
)
|
||||||
|
),
|
||||||
|
Chomp1d(padding),
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
return self.net(x)
|
||||||
|
|
||||||
|
|
||||||
class TemporalBlock(nn.Module):
|
class TemporalBlock(nn.Module):
|
||||||
def __init__(self, n_inputs, n_outputs, kernel_size, stride, dilation, padding, dropout=0.2):
|
def __init__(
|
||||||
super(TemporalBlock, self).__init__()
|
self,
|
||||||
self.conv1 = weight_norm(nn.Conv1d(n_inputs, n_outputs, kernel_size,
|
in_channels: int,
|
||||||
stride=stride, padding=padding, dilation=dilation))
|
out_channels: int,
|
||||||
self.chomp1 = Chomp1d(padding)
|
kernel_size: int,
|
||||||
self.relu1 = nn.ReLU()
|
dilation: int,
|
||||||
|
dropout: float,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.conv1 = CausalConv1d(in_channels, out_channels, kernel_size, dilation=dilation)
|
||||||
|
self.act1 = nn.GELU()
|
||||||
self.dropout1 = nn.Dropout(dropout)
|
self.dropout1 = nn.Dropout(dropout)
|
||||||
|
|
||||||
self.conv2 = weight_norm(nn.Conv1d(n_outputs, n_outputs, kernel_size,
|
self.conv2 = CausalConv1d(out_channels, out_channels, kernel_size, dilation=dilation)
|
||||||
stride=stride, padding=padding, dilation=dilation))
|
self.act2 = nn.GELU()
|
||||||
self.chomp2 = Chomp1d(padding)
|
|
||||||
self.relu2 = nn.ReLU()
|
|
||||||
self.dropout2 = nn.Dropout(dropout)
|
self.dropout2 = nn.Dropout(dropout)
|
||||||
|
|
||||||
self.net = nn.Sequential(self.conv1, self.chomp1, self.relu1, self.dropout1,
|
self.residual = nn.Conv1d(in_channels, out_channels, kernel_size=1) if in_channels != out_channels else nn.Identity()
|
||||||
self.conv2, self.chomp2, self.relu2, self.dropout2)
|
self.final_act = nn.GELU()
|
||||||
self.downsample = nn.Conv1d(n_inputs, n_outputs, 1) if n_inputs != n_outputs else None
|
|
||||||
self.relu = nn.ReLU()
|
|
||||||
self.init_weights()
|
|
||||||
|
|
||||||
def init_weights(self):
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
self.conv1.weight.data.normal_(0, 0.01)
|
residual = self.residual(x)
|
||||||
self.conv2.weight.data.normal_(0, 0.01)
|
out = self.conv1(x)
|
||||||
if self.downsample is not None:
|
out = self.act1(out)
|
||||||
self.downsample.weight.data.normal_(0, 0.01)
|
out = self.dropout1(out)
|
||||||
|
out = self.conv2(out)
|
||||||
|
out = self.act2(out)
|
||||||
|
out = self.dropout2(out)
|
||||||
|
return self.final_act(out + residual)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ReceptiveFieldInfo:
|
||||||
|
receptive_field: int
|
||||||
|
dilations: tuple[int, ...]
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
out = self.net(x)
|
|
||||||
res = x if self.downsample is None else self.downsample(x)
|
|
||||||
return self.relu(out + res)
|
|
||||||
|
|
||||||
class TemporalConvNet(nn.Module):
|
class TemporalConvNet(nn.Module):
|
||||||
def __init__(self, num_inputs, num_channels, kernel_size=2, dropout=0.2):
|
def __init__(self, config: ModelConfig) -> None:
|
||||||
super(TemporalConvNet, self).__init__()
|
super().__init__()
|
||||||
layers = []
|
blocks: list[nn.Module] = []
|
||||||
num_levels = len(num_channels)
|
in_channels = config.input_channels
|
||||||
for i in range(num_levels):
|
dilations: list[int] = []
|
||||||
dilation_size = 2 ** i
|
|
||||||
in_channels = num_inputs if i == 0 else num_channels[i-1]
|
|
||||||
out_channels = num_channels[i]
|
|
||||||
layers += [TemporalBlock(in_channels, out_channels, kernel_size, stride=1, dilation=dilation_size,
|
|
||||||
padding=(kernel_size-1) * dilation_size, dropout=dropout)]
|
|
||||||
|
|
||||||
self.network = nn.Sequential(*layers)
|
for level, out_channels in enumerate(config.tcn_channels):
|
||||||
|
dilation = config.dilation_base ** level
|
||||||
|
dilations.append(dilation)
|
||||||
|
blocks.append(
|
||||||
|
TemporalBlock(
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
kernel_size=config.kernel_size,
|
||||||
|
dilation=dilation,
|
||||||
|
dropout=config.dropout,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
in_channels = out_channels
|
||||||
|
|
||||||
|
self.network = nn.Sequential(*blocks)
|
||||||
|
self.output_head = nn.Conv1d(in_channels, config.output_channels, kernel_size=1)
|
||||||
|
self.receptive_field_info = compute_receptive_field(config)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
features = self.network(x)
|
||||||
|
return self.output_head(features)
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
return self.network(x)
|
|
||||||
|
|
||||||
class BuildingTCN(nn.Module):
|
class BuildingTCN(nn.Module):
|
||||||
def __init__(self, input_size, output_size, num_channels, kernel_size=3, dropout=0.2):
|
def __init__(self, config: ModelConfig) -> None:
|
||||||
super(BuildingTCN, self).__init__()
|
super().__init__()
|
||||||
self.tcn = TemporalConvNet(input_size, num_channels, kernel_size, dropout=dropout)
|
self.config = config
|
||||||
self.linear = nn.Linear(num_channels[-1], output_size)
|
self.tcn = TemporalConvNet(config)
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
# x shape: (batch, seq_len, input_size)
|
if x.ndim != 3:
|
||||||
# TCN needs shape: (batch, input_size, seq_len)
|
raise ValueError(f"Expected input shape (batch, seq, channels), got {tuple(x.shape)}")
|
||||||
x = x.transpose(1, 2)
|
x = x.transpose(1, 2)
|
||||||
y = self.tcn(x)
|
y = self.tcn(x)
|
||||||
# y shape: (batch, num_channels, seq_len)
|
return y.transpose(1, 2)
|
||||||
# linear needs shape: (batch, seq_len, num_channels)
|
|
||||||
y = y.transpose(1, 2)
|
@property
|
||||||
return self.linear(y)
|
def receptive_field(self) -> int:
|
||||||
|
return self.tcn.receptive_field_info.receptive_field
|
||||||
|
|
||||||
|
|
||||||
|
def compute_receptive_field(config: ModelConfig) -> ReceptiveFieldInfo:
|
||||||
|
receptive_field = 1
|
||||||
|
dilations: list[int] = []
|
||||||
|
for level, _ in enumerate(config.tcn_channels):
|
||||||
|
dilation = config.dilation_base ** level
|
||||||
|
dilations.append(dilation)
|
||||||
|
receptive_field += 2 * (config.kernel_size - 1) * dilation
|
||||||
|
return ReceptiveFieldInfo(receptive_field=receptive_field, dilations=tuple(dilations))
|
||||||
|
|
||||||
|
|
||||||
|
def build_model(config: ExperimentConfig | ModelConfig) -> BuildingTCN:
|
||||||
|
model_config = config.model if isinstance(config, ExperimentConfig) else config
|
||||||
|
return BuildingTCN(model_config)
|
||||||
|
|||||||
75
src/sanity_check.py
Normal file
75
src/sanity_check.py
Normal file
@@ -0,0 +1,75 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import matplotlib.pyplot as plt
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from config import make_forward_config
|
||||||
|
from dataset import build_dataloaders, report_to_text
|
||||||
|
|
||||||
|
|
||||||
|
def _to_cpu(tensor: torch.Tensor) -> torch.Tensor:
|
||||||
|
return tensor.detach().cpu()
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
config = make_forward_config()
|
||||||
|
config.data.batch_size = 4
|
||||||
|
config.data.window_size = 1024
|
||||||
|
config.data.window_stride = 256
|
||||||
|
|
||||||
|
loaders, datasets, reports = build_dataloaders(config)
|
||||||
|
train_loader = loaders["train"]
|
||||||
|
train_dataset = datasets["train"]
|
||||||
|
|
||||||
|
print("Dataset loading report:")
|
||||||
|
print(report_to_text(reports))
|
||||||
|
|
||||||
|
batch = next(iter(train_loader))
|
||||||
|
|
||||||
|
time = _to_cpu(batch["time"][0]).numpy()
|
||||||
|
x_norm = batch["x"][0]
|
||||||
|
y_norm = batch["y"][0]
|
||||||
|
|
||||||
|
x = _to_cpu(train_dataset.denormalize_x(x_norm)).squeeze(-1).numpy()
|
||||||
|
y = _to_cpu(train_dataset.denormalize_y(y_norm)).squeeze(-1).numpy()
|
||||||
|
|
||||||
|
if config.data.task == "forward":
|
||||||
|
input_label = f"Input {config.data.base_sensor_code}/{config.data.base_axis}"
|
||||||
|
output_label = f"Output {config.data.response_sensor_code}/{config.data.response_axis}"
|
||||||
|
else:
|
||||||
|
input_label = f"Input {config.data.response_sensor_code}/{config.data.response_axis}"
|
||||||
|
output_label = f"Output {config.data.base_sensor_code}/{config.data.base_axis}"
|
||||||
|
|
||||||
|
fig, axes = plt.subplots(2, 1, figsize=(14, 8), sharex=True)
|
||||||
|
fig.suptitle(f"Alignment sanity check: {batch['file_name'][0]}")
|
||||||
|
|
||||||
|
axes[0].plot(time, x, color="black", linewidth=1.2, label=input_label)
|
||||||
|
axes[0].set_ylabel("Acceleration")
|
||||||
|
axes[0].legend(loc="upper right")
|
||||||
|
axes[0].grid(True, alpha=0.3)
|
||||||
|
|
||||||
|
axes[1].plot(time, y, linewidth=1.1, label=output_label)
|
||||||
|
axes[1].set_xlabel("Time")
|
||||||
|
axes[1].set_ylabel("Acceleration")
|
||||||
|
axes[1].legend(loc="upper right")
|
||||||
|
axes[1].grid(True, alpha=0.3)
|
||||||
|
|
||||||
|
n_guides = min(12, len(time))
|
||||||
|
guide_indices = torch.linspace(0, len(time) - 1, steps=n_guides).round().to(torch.long).tolist()
|
||||||
|
guide_times = [time[i] for i in guide_indices]
|
||||||
|
for ax in axes:
|
||||||
|
for guide_time in guide_times:
|
||||||
|
ax.axvline(guide_time, color="gray", linestyle="--", linewidth=0.8, alpha=0.35)
|
||||||
|
|
||||||
|
plt.tight_layout()
|
||||||
|
|
||||||
|
output_path = Path(__file__).resolve().parents[1] / f"sanity_check_alignment_{config.data.task}.png"
|
||||||
|
plt.savefig(output_path, dpi=180, bbox_inches="tight")
|
||||||
|
print(f"Saved figure to: {output_path}")
|
||||||
|
plt.show()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
675
src/train.py
675
src/train.py
@@ -1,145 +1,586 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import math
|
||||||
|
import random
|
||||||
|
from dataclasses import asdict
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pandas as pd
|
||||||
import torch
|
import torch
|
||||||
import torch.optim as optim
|
import torch.nn.functional as F
|
||||||
from config import *
|
from torch import nn
|
||||||
from dataset import get_dataloaders
|
from torch.optim import AdamW
|
||||||
from model import BuildingTCN
|
from torch.optim.lr_scheduler import ReduceLROnPlateau
|
||||||
|
|
||||||
def masked_l1_loss(pred, target, mask):
|
try:
|
||||||
diff = torch.abs(pred - target) * mask
|
from .config import ExperimentConfig, make_forward_config, make_inverse_config
|
||||||
denom = mask.sum().clamp(min=1.0)
|
from .dataset import build_dataloaders, report_to_text
|
||||||
return diff.sum() / denom
|
from .model import build_model, compute_receptive_field
|
||||||
|
except ImportError:
|
||||||
|
from config import ExperimentConfig, make_forward_config, make_inverse_config
|
||||||
|
from dataset import build_dataloaders, report_to_text
|
||||||
|
from model import build_model, compute_receptive_field
|
||||||
|
|
||||||
def masked_spectral_mag_loss(pred, target, mask):
|
|
||||||
# Apply mask in time domain first so missing labels do not pollute spectrum.
|
|
||||||
pred_masked = pred * mask
|
|
||||||
target_masked = target * mask
|
|
||||||
|
|
||||||
pred_fft = torch.fft.rfft(pred_masked, dim=1)
|
class MixedTimeFrequencyLoss(nn.Module):
|
||||||
target_fft = torch.fft.rfft(target_masked, dim=1)
|
def __init__(
|
||||||
pred_mag = torch.abs(pred_fft)
|
self,
|
||||||
target_mag = torch.abs(target_fft)
|
huber_delta: float,
|
||||||
|
fft_loss_weight: float,
|
||||||
|
fft_loss_min_hz: float,
|
||||||
|
fft_loss_max_hz: float,
|
||||||
|
forward_time_loss_weight: float,
|
||||||
|
forward_rms_loss_weight: float,
|
||||||
|
forward_scale_loss_weight: float,
|
||||||
|
forward_underestimate_loss_weight: float,
|
||||||
|
forward_envelope_loss_weight: float,
|
||||||
|
envelope_kernel_size: int,
|
||||||
|
task: str,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.huber_delta = huber_delta
|
||||||
|
self.fft_loss_weight = fft_loss_weight
|
||||||
|
self.fft_loss_min_hz = fft_loss_min_hz
|
||||||
|
self.fft_loss_max_hz = fft_loss_max_hz
|
||||||
|
self.forward_time_loss_weight = forward_time_loss_weight
|
||||||
|
self.forward_rms_loss_weight = forward_rms_loss_weight
|
||||||
|
self.forward_scale_loss_weight = forward_scale_loss_weight
|
||||||
|
self.forward_underestimate_loss_weight = forward_underestimate_loss_weight
|
||||||
|
self.forward_envelope_loss_weight = forward_envelope_loss_weight
|
||||||
|
self.envelope_kernel_size = envelope_kernel_size
|
||||||
|
self.task = task
|
||||||
|
|
||||||
# Weight each sample/channel by its valid-label ratio.
|
def forward(
|
||||||
valid_ratio = (mask.sum(dim=1) / mask.shape[1]).clamp(min=0.0, max=1.0)
|
self,
|
||||||
freq_weight = valid_ratio.unsqueeze(1).expand_as(pred_mag)
|
prediction: torch.Tensor,
|
||||||
|
target: torch.Tensor,
|
||||||
|
time_tensor: torch.Tensor,
|
||||||
|
) -> dict[str, torch.Tensor]:
|
||||||
|
raw_time_loss = F.huber_loss(prediction, target, delta=self.huber_delta)
|
||||||
|
fft_loss = band_limited_fft_loss(
|
||||||
|
prediction=prediction,
|
||||||
|
target=target,
|
||||||
|
time_tensor=time_tensor,
|
||||||
|
min_hz=self.fft_loss_min_hz,
|
||||||
|
max_hz=self.fft_loss_max_hz,
|
||||||
|
)
|
||||||
|
rms_pred = sequence_rms(prediction)
|
||||||
|
rms_target = sequence_rms(target)
|
||||||
|
rms_loss = F.l1_loss(rms_pred, rms_target) if self.task == "forward" else prediction.new_tensor(0.0)
|
||||||
|
scale_loss = forward_scale_loss(prediction, target) if self.task == "forward" else prediction.new_tensor(0.0)
|
||||||
|
underestimate_loss = (
|
||||||
|
forward_underestimate_loss(prediction, target) if self.task == "forward" else prediction.new_tensor(0.0)
|
||||||
|
)
|
||||||
|
envelope_loss = (
|
||||||
|
forward_envelope_loss(prediction, target, self.envelope_kernel_size)
|
||||||
|
if self.task == "forward"
|
||||||
|
else prediction.new_tensor(0.0)
|
||||||
|
)
|
||||||
|
time_loss = (
|
||||||
|
raw_time_loss * self.forward_time_loss_weight
|
||||||
|
if self.task == "forward"
|
||||||
|
else raw_time_loss
|
||||||
|
)
|
||||||
|
total_loss = (
|
||||||
|
time_loss
|
||||||
|
+ self.fft_loss_weight * fft_loss
|
||||||
|
+ self.forward_rms_loss_weight * rms_loss
|
||||||
|
+ self.forward_scale_loss_weight * scale_loss
|
||||||
|
+ self.forward_underestimate_loss_weight * underestimate_loss
|
||||||
|
+ self.forward_envelope_loss_weight * envelope_loss
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"total": total_loss,
|
||||||
|
"time": raw_time_loss,
|
||||||
|
"weighted_time": time_loss,
|
||||||
|
"fft": fft_loss,
|
||||||
|
"rms_loss": rms_loss,
|
||||||
|
"scale_loss": scale_loss,
|
||||||
|
"underestimate_loss": underestimate_loss,
|
||||||
|
"envelope_loss": envelope_loss,
|
||||||
|
}
|
||||||
|
|
||||||
diff = torch.abs(pred_mag - target_mag) * freq_weight
|
|
||||||
denom = freq_weight.sum().clamp(min=1.0)
|
|
||||||
return diff.sum() / denom
|
|
||||||
|
|
||||||
def masked_corr_loss(pred, target, mask, eps=1e-8):
|
def parse_args() -> argparse.Namespace:
|
||||||
# Compute per-sample/per-channel Pearson correlation with masked timesteps.
|
parser = argparse.ArgumentParser(description="Train single-input single-output TCN for structural response modeling.")
|
||||||
count = mask.sum(dim=1) # (batch, channels)
|
parser.add_argument("--task", choices=("forward", "inverse"), default="forward")
|
||||||
valid = count > 1.0
|
parser.add_argument("--epochs", type=int, default=None)
|
||||||
count_safe = count.clamp(min=1.0)
|
parser.add_argument("--batch-size", type=int, default=None)
|
||||||
|
parser.add_argument("--device", type=str, default=None)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
pred_mean = (pred * mask).sum(dim=1) / count_safe
|
|
||||||
target_mean = (target * mask).sum(dim=1) / count_safe
|
|
||||||
|
|
||||||
pred_centered = (pred - pred_mean.unsqueeze(1)) * mask
|
def make_config(task: str) -> ExperimentConfig:
|
||||||
target_centered = (target - target_mean.unsqueeze(1)) * mask
|
if task == "forward":
|
||||||
|
return make_forward_config()
|
||||||
|
if task == "inverse":
|
||||||
|
return make_inverse_config()
|
||||||
|
raise ValueError(f"Unsupported task: {task}")
|
||||||
|
|
||||||
cov = (pred_centered * target_centered).sum(dim=1)
|
|
||||||
pred_var = (pred_centered ** 2).sum(dim=1)
|
|
||||||
target_var = (target_centered ** 2).sum(dim=1)
|
|
||||||
|
|
||||||
corr = cov / (torch.sqrt(pred_var * target_var) + eps)
|
def set_seed(seed: int) -> None:
|
||||||
corr = torch.where(valid, corr, torch.zeros_like(corr))
|
random.seed(seed)
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
torch.cuda.manual_seed_all(seed)
|
||||||
|
|
||||||
valid_float = valid.float()
|
|
||||||
denom = valid_float.sum().clamp(min=1.0)
|
|
||||||
return ((1.0 - corr) * valid_float).sum() / denom
|
|
||||||
|
|
||||||
def train_model():
|
def resolve_device(device_name: str) -> torch.device:
|
||||||
train_loader, val_loader, _, _, _ = get_dataloaders()
|
if device_name.startswith("cuda") and not torch.cuda.is_available():
|
||||||
|
return torch.device("cpu")
|
||||||
|
return torch.device(device_name)
|
||||||
|
|
||||||
# 初始化前向模型 (1 -> 5)
|
|
||||||
model = BuildingTCN(input_size=1, output_size=5, num_channels=CHANNELS,
|
|
||||||
kernel_size=KERNEL_SIZE, dropout=DROPOUT).to(DEVICE)
|
|
||||||
|
|
||||||
optimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)
|
def move_batch_to_device(batch: dict[str, object], device: torch.device) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=8, factor=0.5)
|
x = batch["x"].to(device)
|
||||||
|
y = batch["y"].to(device)
|
||||||
|
time = batch["time"].to(device)
|
||||||
|
return x, y, time
|
||||||
|
|
||||||
best_val_loss = float('inf')
|
|
||||||
no_improve_epochs = 0
|
|
||||||
early_stop_patience = EARLY_STOP_PATIENCE
|
|
||||||
|
|
||||||
for epoch in range(EPOCHS):
|
def denormalize_prediction(
|
||||||
model.train()
|
prediction_norm: torch.Tensor,
|
||||||
train_loss = 0.0
|
target_norm: torch.Tensor,
|
||||||
train_time_loss = 0.0
|
y_mean: torch.Tensor,
|
||||||
train_spec_loss = 0.0
|
y_std: torch.Tensor,
|
||||||
train_corr_loss = 0.0
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
prediction = prediction_norm * y_std + y_mean
|
||||||
|
target = target_norm * y_std + y_mean
|
||||||
|
return prediction, target
|
||||||
|
|
||||||
for inputs, targets, masks in train_loader:
|
|
||||||
inputs = inputs.to(DEVICE)
|
|
||||||
targets = targets.to(DEVICE)
|
|
||||||
masks = masks.to(DEVICE)
|
|
||||||
|
|
||||||
optimizer.zero_grad()
|
def sequence_rms(signal: torch.Tensor) -> torch.Tensor:
|
||||||
outputs = model(inputs)
|
return torch.sqrt(torch.mean(signal ** 2, dim=1))
|
||||||
time_loss = masked_l1_loss(outputs, targets, masks)
|
|
||||||
spec_loss = masked_spectral_mag_loss(outputs, targets, masks)
|
|
||||||
corr_loss = masked_corr_loss(outputs, targets, masks, eps=CORR_LOSS_EPS)
|
|
||||||
loss = time_loss + SPECTRAL_LOSS_WEIGHT * spec_loss + CORR_LOSS_WEIGHT * corr_loss
|
|
||||||
|
|
||||||
loss.backward()
|
|
||||||
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
|
||||||
optimizer.step()
|
|
||||||
|
|
||||||
train_loss += loss.item()
|
def rms_error(prediction: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
|
||||||
train_time_loss += time_loss.item()
|
pred_rms = sequence_rms(prediction)
|
||||||
train_spec_loss += spec_loss.item()
|
target_rms = sequence_rms(target)
|
||||||
train_corr_loss += corr_loss.item()
|
return torch.abs(pred_rms - target_rms)
|
||||||
|
|
||||||
# 验证
|
|
||||||
model.eval()
|
|
||||||
val_loss = 0.0
|
|
||||||
val_time_loss = 0.0
|
|
||||||
val_spec_loss = 0.0
|
|
||||||
val_corr_loss = 0.0
|
|
||||||
with torch.no_grad():
|
|
||||||
for inputs, targets, masks in val_loader:
|
|
||||||
inputs = inputs.to(DEVICE)
|
|
||||||
targets = targets.to(DEVICE)
|
|
||||||
masks = masks.to(DEVICE)
|
|
||||||
outputs = model(inputs)
|
|
||||||
time_loss = masked_l1_loss(outputs, targets, masks)
|
|
||||||
spec_loss = masked_spectral_mag_loss(outputs, targets, masks)
|
|
||||||
corr_loss = masked_corr_loss(outputs, targets, masks, eps=CORR_LOSS_EPS)
|
|
||||||
loss = time_loss + SPECTRAL_LOSS_WEIGHT * spec_loss + CORR_LOSS_WEIGHT * corr_loss
|
|
||||||
val_loss += loss.item()
|
|
||||||
val_time_loss += time_loss.item()
|
|
||||||
val_spec_loss += spec_loss.item()
|
|
||||||
val_corr_loss += corr_loss.item()
|
|
||||||
|
|
||||||
train_loss /= len(train_loader)
|
def forward_scale_loss(prediction: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
|
||||||
train_time_loss /= len(train_loader)
|
epsilon = 1e-6
|
||||||
train_spec_loss /= len(train_loader)
|
pred_std = prediction.std(dim=1, unbiased=False)
|
||||||
train_corr_loss /= len(train_loader)
|
target_std = target.std(dim=1, unbiased=False)
|
||||||
val_loss /= len(val_loader)
|
pred_peak = prediction.abs().amax(dim=1)
|
||||||
val_time_loss /= len(val_loader)
|
target_peak = target.abs().amax(dim=1)
|
||||||
val_spec_loss /= len(val_loader)
|
pred_rms = sequence_rms(prediction)
|
||||||
val_corr_loss /= len(val_loader)
|
target_rms = sequence_rms(target)
|
||||||
|
std_loss = torch.mean(torch.abs(pred_std - target_std) / (target_std.abs() + epsilon))
|
||||||
|
peak_loss = torch.mean(torch.abs(pred_peak - target_peak) / (target_peak.abs() + epsilon))
|
||||||
|
rms_loss = torch.mean(torch.abs(pred_rms - target_rms) / (target_rms.abs() + epsilon))
|
||||||
|
return (std_loss + peak_loss + rms_loss) / 3.0
|
||||||
|
|
||||||
scheduler.step(val_loss)
|
|
||||||
|
|
||||||
print(
|
def forward_underestimate_loss(prediction: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
|
||||||
f"Epoch {epoch+1}/{EPOCHS} | "
|
epsilon = 1e-6
|
||||||
f"Train Loss: {train_loss:.4f} (L1={train_time_loss:.4f}, Spec={train_spec_loss:.4f}, Corr={train_corr_loss:.4f}) | "
|
pred_rms = sequence_rms(prediction)
|
||||||
f"Val Loss: {val_loss:.4f} (L1={val_time_loss:.4f}, Spec={val_spec_loss:.4f}, Corr={val_corr_loss:.4f})"
|
target_rms = sequence_rms(target)
|
||||||
|
pred_peak = prediction.abs().amax(dim=1)
|
||||||
|
target_peak = target.abs().amax(dim=1)
|
||||||
|
pred_std = prediction.std(dim=1, unbiased=False)
|
||||||
|
target_std = target.std(dim=1, unbiased=False)
|
||||||
|
|
||||||
|
rms_under = torch.relu((target_rms - pred_rms) / (target_rms.abs() + epsilon))
|
||||||
|
peak_under = torch.relu((target_peak - pred_peak) / (target_peak.abs() + epsilon))
|
||||||
|
std_under = torch.relu((target_std - pred_std) / (target_std.abs() + epsilon))
|
||||||
|
return (rms_under.mean() + peak_under.mean() + std_under.mean()) / 3.0
|
||||||
|
|
||||||
|
|
||||||
|
def local_envelope(signal: torch.Tensor, kernel_size: int) -> torch.Tensor:
|
||||||
|
signal_abs = signal.squeeze(-1).abs().unsqueeze(1)
|
||||||
|
effective_kernel = max(3, int(kernel_size))
|
||||||
|
if effective_kernel % 2 == 0:
|
||||||
|
effective_kernel += 1
|
||||||
|
padding = effective_kernel // 2
|
||||||
|
return F.max_pool1d(signal_abs, kernel_size=effective_kernel, stride=1, padding=padding).squeeze(1)
|
||||||
|
|
||||||
|
|
||||||
|
def forward_envelope_loss(prediction: torch.Tensor, target: torch.Tensor, kernel_size: int) -> torch.Tensor:
|
||||||
|
epsilon = 1e-6
|
||||||
|
pred_envelope = local_envelope(prediction, kernel_size)
|
||||||
|
target_envelope = local_envelope(target, kernel_size)
|
||||||
|
relative_error = torch.abs(pred_envelope - target_envelope) / (target_envelope.abs() + epsilon)
|
||||||
|
underestimate_error = torch.relu((target_envelope - pred_envelope) / (target_envelope.abs() + epsilon))
|
||||||
|
return relative_error.mean() + 2.0 * underestimate_error.mean()
|
||||||
|
|
||||||
|
|
||||||
|
def band_limited_fft_loss(
|
||||||
|
prediction: torch.Tensor,
|
||||||
|
target: torch.Tensor,
|
||||||
|
time_tensor: torch.Tensor,
|
||||||
|
min_hz: float,
|
||||||
|
max_hz: float,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
prediction_signal = prediction.squeeze(-1)
|
||||||
|
target_signal = target.squeeze(-1)
|
||||||
|
if prediction_signal.shape[1] < 4:
|
||||||
|
pred_fft = torch.abs(torch.fft.rfft(prediction_signal, dim=1))
|
||||||
|
target_fft = torch.abs(torch.fft.rfft(target_signal, dim=1))
|
||||||
|
return F.mse_loss(pred_fft, target_fft)
|
||||||
|
|
||||||
|
dt = torch.diff(time_tensor, dim=1).mean(dim=1)
|
||||||
|
dt = torch.clamp(dt, min=1e-6)
|
||||||
|
pred_fft = torch.abs(torch.fft.rfft(prediction_signal, dim=1))
|
||||||
|
target_fft = torch.abs(torch.fft.rfft(target_signal, dim=1))
|
||||||
|
fft_bins = pred_fft.shape[1]
|
||||||
|
n = prediction_signal.shape[1]
|
||||||
|
bin_indices = torch.arange(fft_bins, device=prediction.device, dtype=prediction.dtype).unsqueeze(0)
|
||||||
|
frequencies = bin_indices / (dt.unsqueeze(1) * n)
|
||||||
|
band_mask = (frequencies >= min_hz) & (frequencies <= max_hz)
|
||||||
|
band_mask[:, 0] = False
|
||||||
|
|
||||||
|
squared_error = (pred_fft - target_fft) ** 2
|
||||||
|
masked_error = squared_error * band_mask.to(squared_error.dtype)
|
||||||
|
valid_counts = band_mask.sum(dim=1).clamp_min(1).to(squared_error.dtype)
|
||||||
|
per_sample_loss = masked_error.sum(dim=1) / valid_counts
|
||||||
|
fallback_loss = squared_error.mean(dim=1)
|
||||||
|
has_valid_band = band_mask.any(dim=1)
|
||||||
|
per_sample_loss = torch.where(has_valid_band, per_sample_loss, fallback_loss)
|
||||||
|
return per_sample_loss.mean()
|
||||||
|
|
||||||
|
|
||||||
|
def dominant_frequency_error(
|
||||||
|
prediction: torch.Tensor,
|
||||||
|
target: torch.Tensor,
|
||||||
|
time_tensor: torch.Tensor,
|
||||||
|
min_hz: float,
|
||||||
|
max_hz: float,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if prediction.shape[1] < 4:
|
||||||
|
return torch.tensor(0.0, device=prediction.device)
|
||||||
|
|
||||||
|
dt = torch.diff(time_tensor, dim=1).mean(dim=1)
|
||||||
|
dt = torch.clamp(dt, min=1e-6)
|
||||||
|
|
||||||
|
prediction_signal = prediction.squeeze(-1)
|
||||||
|
target_signal = target.squeeze(-1)
|
||||||
|
pred_fft = torch.abs(torch.fft.rfft(prediction_signal, dim=1))
|
||||||
|
target_fft = torch.abs(torch.fft.rfft(target_signal, dim=1))
|
||||||
|
n = prediction_signal.shape[1]
|
||||||
|
fft_bins = pred_fft.shape[1]
|
||||||
|
bin_indices = torch.arange(fft_bins, device=prediction.device, dtype=prediction.dtype).unsqueeze(0)
|
||||||
|
frequencies = bin_indices / (dt.unsqueeze(1) * n)
|
||||||
|
band_mask = (frequencies >= min_hz) & (frequencies <= max_hz)
|
||||||
|
|
||||||
|
if not band_mask.any():
|
||||||
|
return torch.tensor(0.0, device=prediction.device)
|
||||||
|
|
||||||
|
masked_pred_fft = pred_fft.masked_fill(~band_mask, float("-inf"))
|
||||||
|
masked_target_fft = target_fft.masked_fill(~band_mask, float("-inf"))
|
||||||
|
|
||||||
|
pred_index = masked_pred_fft.argmax(dim=1)
|
||||||
|
target_index = masked_target_fft.argmax(dim=1)
|
||||||
|
pred_freq = pred_index.to(prediction.dtype) / (dt * n)
|
||||||
|
target_freq = target_index.to(target.dtype) / (dt * n)
|
||||||
|
return torch.mean(torch.abs(pred_freq - target_freq))
|
||||||
|
|
||||||
|
|
||||||
|
def serialize_for_checkpoint(value: Any) -> Any:
|
||||||
|
if isinstance(value, Path):
|
||||||
|
return str(value)
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return {key: serialize_for_checkpoint(sub_value) for key, sub_value in value.items()}
|
||||||
|
if isinstance(value, tuple):
|
||||||
|
return [serialize_for_checkpoint(item) for item in value]
|
||||||
|
if isinstance(value, list):
|
||||||
|
return [serialize_for_checkpoint(item) for item in value]
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def monitor_metric_name(config: ExperimentConfig) -> str:
|
||||||
|
if config.data.task == "forward":
|
||||||
|
return "rms_error"
|
||||||
|
return "loss"
|
||||||
|
|
||||||
|
|
||||||
|
def run_epoch(
|
||||||
|
model: nn.Module,
|
||||||
|
dataloader: torch.utils.data.DataLoader,
|
||||||
|
optimizer: AdamW | None,
|
||||||
|
criterion: MixedTimeFrequencyLoss,
|
||||||
|
y_mean: torch.Tensor,
|
||||||
|
y_std: torch.Tensor,
|
||||||
|
device: torch.device,
|
||||||
|
grad_clip_norm: float,
|
||||||
|
scaler: torch.cuda.amp.GradScaler,
|
||||||
|
amp_enabled: bool,
|
||||||
|
dominant_freq_min_hz: float,
|
||||||
|
dominant_freq_max_hz: float,
|
||||||
|
) -> dict[str, float]:
|
||||||
|
is_train = optimizer is not None
|
||||||
|
model.train(is_train)
|
||||||
|
|
||||||
|
total_loss_sum = 0.0
|
||||||
|
time_loss_sum = 0.0
|
||||||
|
weighted_time_loss_sum = 0.0
|
||||||
|
fft_loss_sum = 0.0
|
||||||
|
rms_loss_sum = 0.0
|
||||||
|
scale_loss_sum = 0.0
|
||||||
|
underestimate_loss_sum = 0.0
|
||||||
|
envelope_loss_sum = 0.0
|
||||||
|
rms_error_sum = 0.0
|
||||||
|
dominant_freq_error_sum = 0.0
|
||||||
|
sample_count = 0
|
||||||
|
|
||||||
|
for batch in dataloader:
|
||||||
|
x, y_norm, time = move_batch_to_device(batch, device)
|
||||||
|
|
||||||
|
if is_train:
|
||||||
|
optimizer.zero_grad(set_to_none=True)
|
||||||
|
|
||||||
|
with torch.amp.autocast(device_type=device.type, enabled=amp_enabled):
|
||||||
|
prediction_norm = model(x)
|
||||||
|
prediction, target = denormalize_prediction(prediction_norm, y_norm, y_mean, y_std)
|
||||||
|
losses = criterion(prediction, target, time)
|
||||||
|
|
||||||
|
if is_train:
|
||||||
|
scaler.scale(losses["total"]).backward()
|
||||||
|
scaler.unscale_(optimizer)
|
||||||
|
torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm)
|
||||||
|
scaler.step(optimizer)
|
||||||
|
scaler.update()
|
||||||
|
|
||||||
|
batch_size = x.shape[0]
|
||||||
|
total_loss_sum += losses["total"].detach().item() * batch_size
|
||||||
|
time_loss_sum += losses["time"].detach().item() * batch_size
|
||||||
|
weighted_time_loss_sum += losses["weighted_time"].detach().item() * batch_size
|
||||||
|
fft_loss_sum += losses["fft"].detach().item() * batch_size
|
||||||
|
rms_loss_sum += losses["rms_loss"].detach().item() * batch_size
|
||||||
|
scale_loss_sum += losses["scale_loss"].detach().item() * batch_size
|
||||||
|
underestimate_loss_sum += losses["underestimate_loss"].detach().item() * batch_size
|
||||||
|
envelope_loss_sum += losses["envelope_loss"].detach().item() * batch_size
|
||||||
|
rms_error_sum += rms_error(prediction.detach(), target.detach()).mean().item() * batch_size
|
||||||
|
dominant_freq_error_sum += dominant_frequency_error(
|
||||||
|
prediction.detach(),
|
||||||
|
target.detach(),
|
||||||
|
time.detach(),
|
||||||
|
min_hz=dominant_freq_min_hz,
|
||||||
|
max_hz=dominant_freq_max_hz,
|
||||||
|
).item() * batch_size
|
||||||
|
sample_count += batch_size
|
||||||
|
|
||||||
|
if sample_count == 0:
|
||||||
|
raise RuntimeError("No samples were processed in the epoch.")
|
||||||
|
|
||||||
|
return {
|
||||||
|
"loss": total_loss_sum / sample_count,
|
||||||
|
"time_loss": time_loss_sum / sample_count,
|
||||||
|
"weighted_time_loss": weighted_time_loss_sum / sample_count,
|
||||||
|
"fft_loss": fft_loss_sum / sample_count,
|
||||||
|
"rms_loss": rms_loss_sum / sample_count,
|
||||||
|
"scale_loss": scale_loss_sum / sample_count,
|
||||||
|
"underestimate_loss": underestimate_loss_sum / sample_count,
|
||||||
|
"envelope_loss": envelope_loss_sum / sample_count,
|
||||||
|
"rms_error": rms_error_sum / sample_count,
|
||||||
|
"dominant_freq_error": dominant_freq_error_sum / sample_count,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def checkpoint_paths(config: ExperimentConfig) -> tuple[Path, Path]:
|
||||||
|
root = config.data.project_root / config.train.checkpoint_dir / config.data.task
|
||||||
|
root.mkdir(parents=True, exist_ok=True)
|
||||||
|
return root / config.train.best_model_name, root / config.train.history_name
|
||||||
|
|
||||||
|
|
||||||
|
def train_model(config: ExperimentConfig) -> None:
|
||||||
|
set_seed(config.train.seed)
|
||||||
|
device = resolve_device(config.train.device)
|
||||||
|
|
||||||
|
dataloaders, datasets, reports = build_dataloaders(config)
|
||||||
|
train_dataset = datasets["train"]
|
||||||
|
normalization = train_dataset.normalization
|
||||||
|
y_mean = normalization.y_mean.to(device).view(1, 1, -1)
|
||||||
|
y_std = normalization.y_std.to(device).view(1, 1, -1)
|
||||||
|
|
||||||
|
model = build_model(config).to(device)
|
||||||
|
receptive_field = compute_receptive_field(config.model)
|
||||||
|
optimizer = AdamW(
|
||||||
|
model.parameters(),
|
||||||
|
lr=config.train.learning_rate,
|
||||||
|
weight_decay=config.train.weight_decay,
|
||||||
|
)
|
||||||
|
scheduler = ReduceLROnPlateau(
|
||||||
|
optimizer,
|
||||||
|
mode="min",
|
||||||
|
factor=config.train.lr_scheduler_factor,
|
||||||
|
patience=config.train.lr_scheduler_patience,
|
||||||
|
min_lr=config.train.min_learning_rate,
|
||||||
|
)
|
||||||
|
criterion = MixedTimeFrequencyLoss(
|
||||||
|
huber_delta=config.loss.huber_delta,
|
||||||
|
fft_loss_weight=config.loss.fft_loss_weight,
|
||||||
|
fft_loss_min_hz=config.loss.fft_loss_min_hz,
|
||||||
|
fft_loss_max_hz=config.loss.fft_loss_max_hz,
|
||||||
|
forward_time_loss_weight=config.loss.forward_time_loss_weight,
|
||||||
|
forward_rms_loss_weight=config.loss.forward_rms_loss_weight,
|
||||||
|
forward_scale_loss_weight=config.loss.forward_scale_loss_weight,
|
||||||
|
forward_underestimate_loss_weight=config.loss.forward_underestimate_loss_weight,
|
||||||
|
forward_envelope_loss_weight=config.loss.forward_envelope_loss_weight,
|
||||||
|
envelope_kernel_size=config.loss.envelope_kernel_size,
|
||||||
|
task=config.data.task,
|
||||||
|
)
|
||||||
|
|
||||||
|
amp_enabled = config.train.use_amp and device.type == "cuda"
|
||||||
|
scaler = torch.cuda.amp.GradScaler(enabled=amp_enabled)
|
||||||
|
best_model_path, history_path = checkpoint_paths(config)
|
||||||
|
|
||||||
|
print(f"Task: {config.data.task}")
|
||||||
|
print(f"Device: {device}")
|
||||||
|
print(
|
||||||
|
"TCN receptive field: "
|
||||||
|
f"{receptive_field.receptive_field} samples, "
|
||||||
|
f"dilations={receptive_field.dilations}"
|
||||||
|
)
|
||||||
|
print(report_to_text(reports))
|
||||||
|
|
||||||
|
history_rows: list[dict[str, float | int]] = []
|
||||||
|
monitor_name = monitor_metric_name(config)
|
||||||
|
best_monitor_value = math.inf
|
||||||
|
stale_epochs = 0
|
||||||
|
|
||||||
|
for epoch in range(1, config.train.epochs + 1):
|
||||||
|
train_metrics = run_epoch(
|
||||||
|
model=model,
|
||||||
|
dataloader=dataloaders["train"],
|
||||||
|
optimizer=optimizer,
|
||||||
|
criterion=criterion,
|
||||||
|
y_mean=y_mean,
|
||||||
|
y_std=y_std,
|
||||||
|
device=device,
|
||||||
|
grad_clip_norm=config.train.grad_clip_norm,
|
||||||
|
scaler=scaler,
|
||||||
|
amp_enabled=amp_enabled,
|
||||||
|
dominant_freq_min_hz=config.train.dominant_freq_min_hz,
|
||||||
|
dominant_freq_max_hz=config.train.dominant_freq_max_hz,
|
||||||
)
|
)
|
||||||
|
|
||||||
if val_loss < best_val_loss:
|
with torch.no_grad():
|
||||||
best_val_loss = val_loss
|
val_metrics = run_epoch(
|
||||||
torch.save(model.state_dict(), 'best_model.pth')
|
model=model,
|
||||||
print(" --> Saved Best Model")
|
dataloader=dataloaders["val"],
|
||||||
no_improve_epochs = 0
|
optimizer=None,
|
||||||
else:
|
criterion=criterion,
|
||||||
no_improve_epochs += 1
|
y_mean=y_mean,
|
||||||
if ENABLE_EARLY_STOP and no_improve_epochs >= early_stop_patience:
|
y_std=y_std,
|
||||||
print(f"Early stopping at epoch {epoch+1}")
|
device=device,
|
||||||
break
|
grad_clip_norm=config.train.grad_clip_norm,
|
||||||
|
scaler=scaler,
|
||||||
|
amp_enabled=amp_enabled,
|
||||||
|
dominant_freq_min_hz=config.train.dominant_freq_min_hz,
|
||||||
|
dominant_freq_max_hz=config.train.dominant_freq_max_hz,
|
||||||
|
)
|
||||||
|
|
||||||
if __name__ == '__main__':
|
current_monitor_value = val_metrics[monitor_name]
|
||||||
train_model()
|
scheduler.step(current_monitor_value)
|
||||||
|
current_lr = optimizer.param_groups[0]["lr"]
|
||||||
|
|
||||||
|
history_row: dict[str, float | int] = {
|
||||||
|
"epoch": epoch,
|
||||||
|
"lr": current_lr,
|
||||||
|
"train_loss": train_metrics["loss"],
|
||||||
|
"train_time_loss": train_metrics["time_loss"],
|
||||||
|
"train_weighted_time_loss": train_metrics["weighted_time_loss"],
|
||||||
|
"train_fft_loss": train_metrics["fft_loss"],
|
||||||
|
"train_rms_loss": train_metrics["rms_loss"],
|
||||||
|
"train_scale_loss": train_metrics["scale_loss"],
|
||||||
|
"train_underestimate_loss": train_metrics["underestimate_loss"],
|
||||||
|
"train_envelope_loss": train_metrics["envelope_loss"],
|
||||||
|
"train_rms_error": train_metrics["rms_error"],
|
||||||
|
"train_dominant_freq_error": train_metrics["dominant_freq_error"],
|
||||||
|
"val_loss": val_metrics["loss"],
|
||||||
|
"val_time_loss": val_metrics["time_loss"],
|
||||||
|
"val_weighted_time_loss": val_metrics["weighted_time_loss"],
|
||||||
|
"val_fft_loss": val_metrics["fft_loss"],
|
||||||
|
"val_rms_loss": val_metrics["rms_loss"],
|
||||||
|
"val_scale_loss": val_metrics["scale_loss"],
|
||||||
|
"val_underestimate_loss": val_metrics["underestimate_loss"],
|
||||||
|
"val_envelope_loss": val_metrics["envelope_loss"],
|
||||||
|
"val_rms_error": val_metrics["rms_error"],
|
||||||
|
"val_dominant_freq_error": val_metrics["dominant_freq_error"],
|
||||||
|
}
|
||||||
|
history_rows.append(history_row)
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"Epoch {epoch:03d} | "
|
||||||
|
f"train_loss={train_metrics['loss']:.6f} | "
|
||||||
|
f"val_loss={val_metrics['loss']:.6f} | "
|
||||||
|
f"val_rms={val_metrics['rms_error']:.6f} | "
|
||||||
|
f"val_dom_freq_err={val_metrics['dominant_freq_error']:.6f} Hz | "
|
||||||
|
f"monitor=val_{monitor_name}:{current_monitor_value:.6f} | "
|
||||||
|
f"lr={current_lr:.2e}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if current_monitor_value < best_monitor_value:
|
||||||
|
best_monitor_value = current_monitor_value
|
||||||
|
stale_epochs = 0
|
||||||
|
torch.save(
|
||||||
|
{
|
||||||
|
"model_state_dict": model.state_dict(),
|
||||||
|
"optimizer_state_dict": optimizer.state_dict(),
|
||||||
|
"config": serialize_for_checkpoint(asdict(config)),
|
||||||
|
"epoch": epoch,
|
||||||
|
"best_monitor_name": monitor_name,
|
||||||
|
"best_monitor_value": best_monitor_value,
|
||||||
|
"val_metrics": val_metrics,
|
||||||
|
},
|
||||||
|
best_model_path,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
stale_epochs += 1
|
||||||
|
|
||||||
|
if stale_epochs >= config.train.early_stop_patience:
|
||||||
|
print(f"Early stopping triggered after {epoch} epochs.")
|
||||||
|
break
|
||||||
|
|
||||||
|
history_df = pd.DataFrame(history_rows)
|
||||||
|
history_df.to_csv(history_path, index=False)
|
||||||
|
print(f"Best model saved to: {best_model_path}")
|
||||||
|
print(f"Training history saved to: {history_path}")
|
||||||
|
|
||||||
|
if best_model_path.exists():
|
||||||
|
checkpoint = torch.load(best_model_path, map_location=device, weights_only=False)
|
||||||
|
model.load_state_dict(checkpoint["model_state_dict"])
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
test_metrics = run_epoch(
|
||||||
|
model=model,
|
||||||
|
dataloader=dataloaders["test"],
|
||||||
|
optimizer=None,
|
||||||
|
criterion=criterion,
|
||||||
|
y_mean=y_mean,
|
||||||
|
y_std=y_std,
|
||||||
|
device=device,
|
||||||
|
grad_clip_norm=config.train.grad_clip_norm,
|
||||||
|
scaler=scaler,
|
||||||
|
amp_enabled=amp_enabled,
|
||||||
|
dominant_freq_min_hz=config.train.dominant_freq_min_hz,
|
||||||
|
dominant_freq_max_hz=config.train.dominant_freq_max_hz,
|
||||||
|
)
|
||||||
|
|
||||||
|
print(
|
||||||
|
"Test metrics | "
|
||||||
|
f"loss={test_metrics['loss']:.6f} | "
|
||||||
|
f"rms={test_metrics['rms_error']:.6f} | "
|
||||||
|
f"dominant_freq_error={test_metrics['dominant_freq_error']:.6f} Hz"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
args = parse_args()
|
||||||
|
config = make_config(args.task)
|
||||||
|
|
||||||
|
if args.epochs is not None:
|
||||||
|
config.train.epochs = args.epochs
|
||||||
|
if args.batch_size is not None:
|
||||||
|
config.data.batch_size = args.batch_size
|
||||||
|
if args.device is not None:
|
||||||
|
config.train.device = args.device
|
||||||
|
|
||||||
|
train_model(config)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
|
|||||||
BIN
src_old/__pycache__/config.cpython-310.pyc
Normal file
BIN
src_old/__pycache__/config.cpython-310.pyc
Normal file
Binary file not shown.
BIN
src_old/__pycache__/config.cpython-314.pyc
Normal file
BIN
src_old/__pycache__/config.cpython-314.pyc
Normal file
Binary file not shown.
BIN
src_old/__pycache__/dataset.cpython-310.pyc
Normal file
BIN
src_old/__pycache__/dataset.cpython-310.pyc
Normal file
Binary file not shown.
BIN
src_old/__pycache__/dataset.cpython-314.pyc
Normal file
BIN
src_old/__pycache__/dataset.cpython-314.pyc
Normal file
Binary file not shown.
BIN
src_old/__pycache__/evaluate.cpython-314.pyc
Normal file
BIN
src_old/__pycache__/evaluate.cpython-314.pyc
Normal file
Binary file not shown.
BIN
src_old/__pycache__/model.cpython-310.pyc
Normal file
BIN
src_old/__pycache__/model.cpython-310.pyc
Normal file
Binary file not shown.
BIN
src_old/__pycache__/model.cpython-314.pyc
Normal file
BIN
src_old/__pycache__/model.cpython-314.pyc
Normal file
Binary file not shown.
BIN
src_old/__pycache__/train.cpython-310.pyc
Normal file
BIN
src_old/__pycache__/train.cpython-310.pyc
Normal file
Binary file not shown.
BIN
src_old/__pycache__/train.cpython-314.pyc
Normal file
BIN
src_old/__pycache__/train.cpython-314.pyc
Normal file
Binary file not shown.
36
src_old/config.py
Normal file
36
src_old/config.py
Normal file
@@ -0,0 +1,36 @@
|
|||||||
|
import os
|
||||||
|
|
||||||
|
# Data Configuration
|
||||||
|
# 使用基于当前文件的绝对路径拼接,以防止你在不同目录下运行报错
|
||||||
|
DATA_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), 'downloads')
|
||||||
|
SEQ_LEN = 512 # 滑动窗口的长度 (时间序列步数)
|
||||||
|
STEP_SIZE = 20 # 滑动窗口的步长
|
||||||
|
BATCH_SIZE = 128
|
||||||
|
|
||||||
|
# Features Configuration
|
||||||
|
INPUT_SENSOR = 'WSMS00012'
|
||||||
|
OUTPUT_SENSORS = ['WSMS00007', 'WSMS00008', 'WSMS00009', 'WSMS00010', 'WSMS00011']
|
||||||
|
INPUT_AXIS = 'value1' # 底部传感器输入轴(课程要求:X轴)
|
||||||
|
OUTPUT_AXIS = 'value3' # 目标传感器输出轴
|
||||||
|
|
||||||
|
# Model Configuration
|
||||||
|
CHANNELS = [64, 64, 128, 128, 256, 256] # TCN 各层通道数
|
||||||
|
KERNEL_SIZE = 5
|
||||||
|
DROPOUT = 0.1
|
||||||
|
|
||||||
|
# Training Configuration
|
||||||
|
LEARNING_RATE = 3e-4
|
||||||
|
EPOCHS = 50
|
||||||
|
WEIGHT_DECAY = 1e-4
|
||||||
|
ENABLE_EARLY_STOP = False
|
||||||
|
EARLY_STOP_PATIENCE = 15
|
||||||
|
SPECTRAL_LOSS_WEIGHT = 0.25
|
||||||
|
CORR_LOSS_WEIGHT = 0.05
|
||||||
|
MSE_LOSS_WEIGHT = 0.5
|
||||||
|
STD_LOSS_WEIGHT = 0.3
|
||||||
|
PEAK_LOSS_WEIGHT = 0.2
|
||||||
|
CORR_LOSS_EPS = 1e-8
|
||||||
|
|
||||||
|
# Device
|
||||||
|
import torch
|
||||||
|
DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||||
188
src_old/dataset.py
Normal file
188
src_old/dataset.py
Normal file
@@ -0,0 +1,188 @@
|
|||||||
|
import os
|
||||||
|
import glob
|
||||||
|
import pandas as pd
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from torch.utils.data import Dataset, DataLoader
|
||||||
|
from sklearn.preprocessing import StandardScaler
|
||||||
|
from config import *
|
||||||
|
|
||||||
|
class MultiOutputStandardizer:
|
||||||
|
"""Per-channel standardization that ignores missing labels via masks."""
|
||||||
|
def __init__(self, n_outputs):
|
||||||
|
self.n_outputs = n_outputs
|
||||||
|
self.mean_ = np.zeros(n_outputs, dtype=np.float32)
|
||||||
|
self.scale_ = np.ones(n_outputs, dtype=np.float32)
|
||||||
|
self.fitted = False
|
||||||
|
|
||||||
|
def fit(self, y_sequences, mask_sequences):
|
||||||
|
means = []
|
||||||
|
scales = []
|
||||||
|
for c in range(self.n_outputs):
|
||||||
|
valid_values = []
|
||||||
|
for y_seq, m_seq in zip(y_sequences, mask_sequences):
|
||||||
|
valid = m_seq[:, c] > 0.5
|
||||||
|
if np.any(valid):
|
||||||
|
valid_values.append(y_seq[valid, c])
|
||||||
|
if len(valid_values) == 0:
|
||||||
|
means.append(0.0)
|
||||||
|
scales.append(1.0)
|
||||||
|
continue
|
||||||
|
vals = np.concatenate(valid_values, axis=0)
|
||||||
|
mean = float(np.mean(vals))
|
||||||
|
std = float(np.std(vals))
|
||||||
|
if std < 1e-6:
|
||||||
|
std = 1.0
|
||||||
|
means.append(mean)
|
||||||
|
scales.append(std)
|
||||||
|
|
||||||
|
self.mean_ = np.asarray(means, dtype=np.float32)
|
||||||
|
self.scale_ = np.asarray(scales, dtype=np.float32)
|
||||||
|
self.fitted = True
|
||||||
|
|
||||||
|
def transform(self, y):
|
||||||
|
if not self.fitted:
|
||||||
|
raise RuntimeError("MultiOutputStandardizer must be fitted before transform.")
|
||||||
|
return (y - self.mean_) / self.scale_
|
||||||
|
|
||||||
|
def inverse_transform(self, y):
|
||||||
|
if not self.fitted:
|
||||||
|
raise RuntimeError("MultiOutputStandardizer must be fitted before inverse_transform.")
|
||||||
|
return y * self.scale_ + self.mean_
|
||||||
|
|
||||||
|
|
||||||
|
class BuildingDataset(Dataset):
|
||||||
|
def __init__(self, file_paths, seq_len, step_size, scaler_X=None, scaler_Y=None, fit_scaler=False):
|
||||||
|
self.seq_len = seq_len
|
||||||
|
self.X_data = []
|
||||||
|
self.Y_data = []
|
||||||
|
self.M_data = []
|
||||||
|
self.scaler_X = scaler_X if scaler_X is not None else StandardScaler()
|
||||||
|
self.scaler_Y = scaler_Y if scaler_Y is not None else MultiOutputStandardizer(len(OUTPUT_SENSORS))
|
||||||
|
|
||||||
|
raw_X = []
|
||||||
|
raw_Y = []
|
||||||
|
raw_M = []
|
||||||
|
|
||||||
|
for f in file_paths:
|
||||||
|
# 读取数据
|
||||||
|
df = pd.read_csv(f)
|
||||||
|
# 使用长表格式: code, type, time, value1, value2, value3
|
||||||
|
|
||||||
|
# 提取 012 的输入轴作为基准
|
||||||
|
df_in = df[df['code'] == INPUT_SENSOR][['time', INPUT_AXIS]].rename(columns={INPUT_AXIS: 'input_signal'})
|
||||||
|
if len(df_in) == 0:
|
||||||
|
# 若无输入传感器(如自由衰减数据),则补零
|
||||||
|
df_in = pd.DataFrame({'time': df['time'].unique()})
|
||||||
|
df_in['input_signal'] = 0.0
|
||||||
|
|
||||||
|
# 提取 007~011 的 Z 轴并按时间戳逐步合并 (使用 left join 确保以 df_in 的时间为基准)
|
||||||
|
df_merged = df_in
|
||||||
|
for sens in OUTPUT_SENSORS:
|
||||||
|
df_out_sens = df[df['code'] == sens][['time', OUTPUT_AXIS]].rename(columns={OUTPUT_AXIS: f'out_{sens}'})
|
||||||
|
df_merged = pd.merge(df_merged, df_out_sens, on='time', how='left')
|
||||||
|
|
||||||
|
df_merged = df_merged.sort_values('time').reset_index(drop=True)
|
||||||
|
|
||||||
|
if len(df_merged) == 0:
|
||||||
|
print(f"Warning: Skipping file {f} due to no overlapping timestamps across required sensors.")
|
||||||
|
continue
|
||||||
|
|
||||||
|
x_seq = df_merged['input_signal'].values.reshape(-1, 1).astype(np.float32)
|
||||||
|
|
||||||
|
# 提取所有 target 传感器列与可用性掩码
|
||||||
|
out_cols = [f'out_{sens}' for sens in OUTPUT_SENSORS]
|
||||||
|
y_seq = np.zeros((len(df_merged), len(OUTPUT_SENSORS)), dtype=np.float32)
|
||||||
|
m_seq = np.zeros((len(df_merged), len(OUTPUT_SENSORS)), dtype=np.float32)
|
||||||
|
for c, col in enumerate(out_cols):
|
||||||
|
series = df_merged[col]
|
||||||
|
observed = ~series.isna()
|
||||||
|
m_seq[:, c] = observed.astype(np.float32)
|
||||||
|
if observed.any():
|
||||||
|
filled = series.interpolate(method='linear').bfill().ffill()
|
||||||
|
y_seq[:, c] = filled.fillna(0.0).values.astype(np.float32)
|
||||||
|
else:
|
||||||
|
y_seq[:, c] = 0.0
|
||||||
|
|
||||||
|
raw_X.append(x_seq)
|
||||||
|
raw_Y.append(y_seq)
|
||||||
|
raw_M.append(m_seq)
|
||||||
|
|
||||||
|
if len(raw_X) == 0:
|
||||||
|
raise ValueError("未能从文件中构造出有效序列,请检查数据路径与传感器编码配置。")
|
||||||
|
|
||||||
|
# 拼接所有文件数据进行 fit
|
||||||
|
X_all = np.vstack(raw_X)
|
||||||
|
|
||||||
|
if fit_scaler:
|
||||||
|
self.scaler_X.fit(X_all)
|
||||||
|
self.scaler_Y.fit(raw_Y, raw_M)
|
||||||
|
|
||||||
|
# 切分窗口
|
||||||
|
for x_seq, y_seq, m_seq in zip(raw_X, raw_Y, raw_M):
|
||||||
|
x_seq_scaled = self.scaler_X.transform(x_seq)
|
||||||
|
y_seq_scaled = self.scaler_Y.transform(y_seq)
|
||||||
|
y_seq_scaled = np.where(m_seq > 0.5, y_seq_scaled, 0.0).astype(np.float32)
|
||||||
|
|
||||||
|
for i in range(0, len(x_seq_scaled) - seq_len + 1, step_size):
|
||||||
|
x_win = x_seq_scaled[i:i+seq_len]
|
||||||
|
y_win = y_seq_scaled[i:i+seq_len]
|
||||||
|
m_win = m_seq[i:i+seq_len]
|
||||||
|
if np.sum(m_win) <= 0:
|
||||||
|
continue
|
||||||
|
self.X_data.append(x_win)
|
||||||
|
self.Y_data.append(y_win)
|
||||||
|
self.M_data.append(m_win)
|
||||||
|
|
||||||
|
self.X_data = np.array(self.X_data)
|
||||||
|
self.Y_data = np.array(self.Y_data)
|
||||||
|
self.M_data = np.array(self.M_data)
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.X_data)
|
||||||
|
|
||||||
|
def __getitem__(self, idx):
|
||||||
|
return (
|
||||||
|
torch.tensor(self.X_data[idx], dtype=torch.float32),
|
||||||
|
torch.tensor(self.Y_data[idx], dtype=torch.float32),
|
||||||
|
torch.tensor(self.M_data[idx], dtype=torch.float32),
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_dataloaders(condition='Non_TMD', include_free_vib=False):
|
||||||
|
"""
|
||||||
|
condition: 'Non_TMD' 或者是 'TMD'
|
||||||
|
include_free_vib: 是否在训练集中加入自由振动与自由衰减数据
|
||||||
|
"""
|
||||||
|
base_dir = os.path.join(DATA_DIR, condition)
|
||||||
|
|
||||||
|
# 手动区分的子目录
|
||||||
|
train_files = glob.glob(os.path.join(base_dir, 'train', '*.csv'))
|
||||||
|
if include_free_vib:
|
||||||
|
train_files += glob.glob(os.path.join(base_dir, 'free_vib', '*.csv'))
|
||||||
|
|
||||||
|
val_files = glob.glob(os.path.join(base_dir, 'val', '*.csv'))
|
||||||
|
test_files = glob.glob(os.path.join(base_dir, 'test', '*.csv'))
|
||||||
|
|
||||||
|
print(f"[{condition}] Train files: {len(train_files)}, Val files: {len(val_files)}, Test files: {len(test_files)}")
|
||||||
|
|
||||||
|
if len(train_files) == 0:
|
||||||
|
raise ValueError(f"错误: 在 {base_dir}/train 目录下未找到训练文件!请检查路径是否正确。")
|
||||||
|
if len(val_files) == 0:
|
||||||
|
raise ValueError(f"错误: 在 {base_dir}/val 目录下未找到验证文件!请检查路径是否正确。")
|
||||||
|
|
||||||
|
train_dataset = BuildingDataset(train_files, SEQ_LEN, STEP_SIZE, fit_scaler=True)
|
||||||
|
val_dataset = BuildingDataset(val_files, SEQ_LEN, STEP_SIZE,
|
||||||
|
scaler_X=train_dataset.scaler_X,
|
||||||
|
scaler_Y=train_dataset.scaler_Y, fit_scaler=False)
|
||||||
|
# 若某条件(如 TMD)下没有 test 数据,可以处理一下防止报错
|
||||||
|
test_loader = None
|
||||||
|
if len(test_files) > 0:
|
||||||
|
test_dataset = BuildingDataset(test_files, SEQ_LEN, STEP_SIZE,
|
||||||
|
scaler_X=train_dataset.scaler_X,
|
||||||
|
scaler_Y=train_dataset.scaler_Y, fit_scaler=False)
|
||||||
|
test_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False)
|
||||||
|
|
||||||
|
train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)
|
||||||
|
val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False)
|
||||||
|
|
||||||
|
return train_loader, val_loader, test_loader, train_dataset.scaler_X, train_dataset.scaler_Y
|
||||||
73
src_old/evaluate.py
Normal file
73
src_old/evaluate.py
Normal file
@@ -0,0 +1,73 @@
|
|||||||
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
import matplotlib.pyplot as plt
|
||||||
|
from config import *
|
||||||
|
from dataset import get_dataloaders
|
||||||
|
from model import BuildingTCN
|
||||||
|
|
||||||
|
def evaluate_model():
|
||||||
|
_, _, test_loader, scaler_X, scaler_Y = get_dataloaders()
|
||||||
|
if test_loader is None:
|
||||||
|
raise ValueError("当前数据配置下没有可用的测试集。")
|
||||||
|
|
||||||
|
model = BuildingTCN(input_size=1, output_size=5, num_channels=CHANNELS,
|
||||||
|
kernel_size=KERNEL_SIZE, dropout=DROPOUT).to(DEVICE)
|
||||||
|
model.load_state_dict(torch.load('best_model.pth', map_location=DEVICE))
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
all_preds = []
|
||||||
|
all_targets = []
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for inputs, targets, masks in test_loader:
|
||||||
|
inputs = inputs.to(DEVICE)
|
||||||
|
outputs = model(inputs)
|
||||||
|
|
||||||
|
all_preds.append(outputs.cpu().numpy())
|
||||||
|
all_targets.append(targets.numpy())
|
||||||
|
|
||||||
|
# Concatenate results
|
||||||
|
all_preds = np.concatenate(all_preds, axis=0)
|
||||||
|
all_targets = np.concatenate(all_targets, axis=0)
|
||||||
|
|
||||||
|
# Inverse transform to physical scale for metrics
|
||||||
|
all_preds_inv = scaler_Y.inverse_transform(all_preds.reshape(-1, len(OUTPUT_SENSORS))).reshape(all_preds.shape)
|
||||||
|
all_targets_inv = scaler_Y.inverse_transform(all_targets.reshape(-1, len(OUTPUT_SENSORS))).reshape(all_targets.shape)
|
||||||
|
|
||||||
|
print("\n=== Test Metrics (All Windows) ===")
|
||||||
|
for i, sens in enumerate(OUTPUT_SENSORS):
|
||||||
|
pred_i = all_preds_inv[:, :, i].reshape(-1)
|
||||||
|
target_i = all_targets_inv[:, :, i].reshape(-1)
|
||||||
|
mae = float(np.mean(np.abs(pred_i - target_i)))
|
||||||
|
pred_std = float(np.std(pred_i))
|
||||||
|
target_std = float(np.std(target_i))
|
||||||
|
amp_ratio = pred_std / (target_std + 1e-12)
|
||||||
|
pred_peak = float(np.max(np.abs(pred_i)))
|
||||||
|
target_peak = float(np.max(np.abs(target_i)))
|
||||||
|
peak_ratio = pred_peak / (target_peak + 1e-12)
|
||||||
|
corr = float(np.corrcoef(pred_i, target_i)[0, 1]) if len(pred_i) > 1 else float("nan")
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"{sens}: MAE={mae:.4f}, Corr={corr:.4f}, "
|
||||||
|
f"AmpRatio(std_pred/std_true)={amp_ratio:.4f}, "
|
||||||
|
f"PeakRatio(max|pred|/max|true|)={peak_ratio:.4f}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 取一个 batch 的第一条序列进行可视化
|
||||||
|
sample_pred = all_preds_inv[0]
|
||||||
|
sample_target = all_targets_inv[0]
|
||||||
|
|
||||||
|
plt.figure(figsize=(15, 10))
|
||||||
|
for i, sens in enumerate(OUTPUT_SENSORS):
|
||||||
|
plt.subplot(5, 1, i+1)
|
||||||
|
plt.plot(sample_target[:, i], label='True', alpha=0.7)
|
||||||
|
plt.plot(sample_pred[:, i], label='Pred', alpha=0.7, linestyle='--')
|
||||||
|
plt.title(f'Sensor {sens} Z-axis Response (Test: Earthquake)')
|
||||||
|
plt.legend()
|
||||||
|
|
||||||
|
plt.tight_layout()
|
||||||
|
plt.savefig('test_results.png')
|
||||||
|
print("Evaluation done. Result saved to test_results.png")
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
evaluate_model()
|
||||||
76
src_old/model.py
Normal file
76
src_old/model.py
Normal file
@@ -0,0 +1,76 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch.nn.utils import weight_norm
|
||||||
|
|
||||||
|
class Chomp1d(nn.Module):
|
||||||
|
def __init__(self, chomp_size):
|
||||||
|
super(Chomp1d, self).__init__()
|
||||||
|
self.chomp_size = chomp_size
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
return x[:, :, :-self.chomp_size].contiguous()
|
||||||
|
|
||||||
|
class TemporalBlock(nn.Module):
|
||||||
|
def __init__(self, n_inputs, n_outputs, kernel_size, stride, dilation, padding, dropout=0.2):
|
||||||
|
super(TemporalBlock, self).__init__()
|
||||||
|
self.conv1 = weight_norm(nn.Conv1d(n_inputs, n_outputs, kernel_size,
|
||||||
|
stride=stride, padding=padding, dilation=dilation))
|
||||||
|
self.chomp1 = Chomp1d(padding)
|
||||||
|
self.relu1 = nn.ReLU()
|
||||||
|
self.dropout1 = nn.Dropout(dropout)
|
||||||
|
|
||||||
|
self.conv2 = weight_norm(nn.Conv1d(n_outputs, n_outputs, kernel_size,
|
||||||
|
stride=stride, padding=padding, dilation=dilation))
|
||||||
|
self.chomp2 = Chomp1d(padding)
|
||||||
|
self.relu2 = nn.ReLU()
|
||||||
|
self.dropout2 = nn.Dropout(dropout)
|
||||||
|
|
||||||
|
self.net = nn.Sequential(self.conv1, self.chomp1, self.relu1, self.dropout1,
|
||||||
|
self.conv2, self.chomp2, self.relu2, self.dropout2)
|
||||||
|
self.downsample = nn.Conv1d(n_inputs, n_outputs, 1) if n_inputs != n_outputs else None
|
||||||
|
self.relu = nn.ReLU()
|
||||||
|
self.init_weights()
|
||||||
|
|
||||||
|
def init_weights(self):
|
||||||
|
self.conv1.weight.data.normal_(0, 0.01)
|
||||||
|
self.conv2.weight.data.normal_(0, 0.01)
|
||||||
|
if self.downsample is not None:
|
||||||
|
self.downsample.weight.data.normal_(0, 0.01)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
out = self.net(x)
|
||||||
|
res = x if self.downsample is None else self.downsample(x)
|
||||||
|
return self.relu(out + res)
|
||||||
|
|
||||||
|
class TemporalConvNet(nn.Module):
|
||||||
|
def __init__(self, num_inputs, num_channels, kernel_size=2, dropout=0.2):
|
||||||
|
super(TemporalConvNet, self).__init__()
|
||||||
|
layers = []
|
||||||
|
num_levels = len(num_channels)
|
||||||
|
for i in range(num_levels):
|
||||||
|
dilation_size = 2 ** i
|
||||||
|
in_channels = num_inputs if i == 0 else num_channels[i-1]
|
||||||
|
out_channels = num_channels[i]
|
||||||
|
layers += [TemporalBlock(in_channels, out_channels, kernel_size, stride=1, dilation=dilation_size,
|
||||||
|
padding=(kernel_size-1) * dilation_size, dropout=dropout)]
|
||||||
|
|
||||||
|
self.network = nn.Sequential(*layers)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
return self.network(x)
|
||||||
|
|
||||||
|
class BuildingTCN(nn.Module):
|
||||||
|
def __init__(self, input_size, output_size, num_channels, kernel_size=3, dropout=0.2):
|
||||||
|
super(BuildingTCN, self).__init__()
|
||||||
|
self.tcn = TemporalConvNet(input_size, num_channels, kernel_size, dropout=dropout)
|
||||||
|
self.linear = nn.Linear(num_channels[-1], output_size)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
# x shape: (batch, seq_len, input_size)
|
||||||
|
# TCN needs shape: (batch, input_size, seq_len)
|
||||||
|
x = x.transpose(1, 2)
|
||||||
|
y = self.tcn(x)
|
||||||
|
# y shape: (batch, num_channels, seq_len)
|
||||||
|
# linear needs shape: (batch, seq_len, num_channels)
|
||||||
|
y = y.transpose(1, 2)
|
||||||
|
return self.linear(y)
|
||||||
241
src_old/train.py
Normal file
241
src_old/train.py
Normal file
@@ -0,0 +1,241 @@
|
|||||||
|
import torch
|
||||||
|
import torch.optim as optim
|
||||||
|
from config import *
|
||||||
|
from dataset import get_dataloaders
|
||||||
|
from model import BuildingTCN
|
||||||
|
|
||||||
|
def masked_l1_loss(pred, target, mask):
|
||||||
|
diff = torch.abs(pred - target) * mask
|
||||||
|
denom = mask.sum().clamp(min=1.0)
|
||||||
|
return diff.sum() / denom
|
||||||
|
|
||||||
|
def masked_mse_loss(pred, target, mask):
|
||||||
|
sq = ((pred - target) ** 2) * mask
|
||||||
|
denom = mask.sum().clamp(min=1.0)
|
||||||
|
return sq.sum() / denom
|
||||||
|
|
||||||
|
def masked_spectral_mag_loss(pred, target, mask):
|
||||||
|
# Apply mask in time domain first so missing labels do not pollute spectrum.
|
||||||
|
pred_masked = pred * mask
|
||||||
|
target_masked = target * mask
|
||||||
|
|
||||||
|
pred_fft = torch.fft.rfft(pred_masked, dim=1)
|
||||||
|
target_fft = torch.fft.rfft(target_masked, dim=1)
|
||||||
|
pred_mag = torch.abs(pred_fft)
|
||||||
|
target_mag = torch.abs(target_fft)
|
||||||
|
pred_mag = torch.nan_to_num(pred_mag, nan=0.0, posinf=1e6, neginf=0.0)
|
||||||
|
target_mag = torch.nan_to_num(target_mag, nan=0.0, posinf=1e6, neginf=0.0)
|
||||||
|
|
||||||
|
# Weight each sample/channel by its valid-label ratio.
|
||||||
|
valid_ratio = (mask.sum(dim=1) / mask.shape[1]).clamp(min=0.0, max=1.0)
|
||||||
|
freq_weight = valid_ratio.unsqueeze(1).expand_as(pred_mag)
|
||||||
|
|
||||||
|
diff = torch.abs(pred_mag - target_mag) * freq_weight
|
||||||
|
denom = freq_weight.sum().clamp(min=1.0)
|
||||||
|
return diff.sum() / denom
|
||||||
|
|
||||||
|
def masked_corr_loss(pred, target, mask, eps=1e-8):
|
||||||
|
# Compute per-sample/per-channel Pearson correlation with masked timesteps.
|
||||||
|
count = mask.sum(dim=1) # (batch, channels)
|
||||||
|
valid = count > 1.0
|
||||||
|
count_safe = count.clamp(min=1.0)
|
||||||
|
|
||||||
|
pred_mean = (pred * mask).sum(dim=1) / count_safe
|
||||||
|
target_mean = (target * mask).sum(dim=1) / count_safe
|
||||||
|
|
||||||
|
pred_centered = (pred - pred_mean.unsqueeze(1)) * mask
|
||||||
|
target_centered = (target - target_mean.unsqueeze(1)) * mask
|
||||||
|
|
||||||
|
cov = (pred_centered * target_centered).sum(dim=1)
|
||||||
|
pred_var = (pred_centered ** 2).sum(dim=1)
|
||||||
|
target_var = (target_centered ** 2).sum(dim=1)
|
||||||
|
|
||||||
|
# Put eps inside sqrt to avoid infinite gradients around zero variance.
|
||||||
|
denom = torch.sqrt((pred_var * target_var).clamp(min=0.0) + eps)
|
||||||
|
corr = cov / denom
|
||||||
|
corr = torch.nan_to_num(corr, nan=0.0, posinf=0.0, neginf=0.0)
|
||||||
|
corr = torch.where(valid, corr, torch.zeros_like(corr))
|
||||||
|
|
||||||
|
valid_float = valid.float()
|
||||||
|
denom = valid_float.sum().clamp(min=1.0)
|
||||||
|
return ((1.0 - corr) * valid_float).sum() / denom
|
||||||
|
|
||||||
|
def masked_std_loss(pred, target, mask):
|
||||||
|
count = mask.sum(dim=1).clamp(min=1.0)
|
||||||
|
pred_mean = (pred * mask).sum(dim=1) / count
|
||||||
|
target_mean = (target * mask).sum(dim=1) / count
|
||||||
|
|
||||||
|
pred_centered = (pred - pred_mean.unsqueeze(1)) * mask
|
||||||
|
target_centered = (target - target_mean.unsqueeze(1)) * mask
|
||||||
|
|
||||||
|
pred_std = torch.sqrt((pred_centered ** 2).sum(dim=1) / count + 1e-8)
|
||||||
|
target_std = torch.sqrt((target_centered ** 2).sum(dim=1) / count + 1e-8)
|
||||||
|
|
||||||
|
valid = mask.sum(dim=1) > 1.0
|
||||||
|
valid_float = valid.float()
|
||||||
|
denom = valid_float.sum().clamp(min=1.0)
|
||||||
|
return (torch.abs(pred_std - target_std) * valid_float).sum() / denom
|
||||||
|
|
||||||
|
def masked_peak_loss(pred, target, mask):
|
||||||
|
abs_pred = torch.abs(pred) * mask
|
||||||
|
abs_target = torch.abs(target) * mask
|
||||||
|
pred_peak = abs_pred.max(dim=1).values
|
||||||
|
target_peak = abs_target.max(dim=1).values
|
||||||
|
|
||||||
|
valid = mask.sum(dim=1) > 0.0
|
||||||
|
valid_float = valid.float()
|
||||||
|
denom = valid_float.sum().clamp(min=1.0)
|
||||||
|
return (torch.abs(pred_peak - target_peak) * valid_float).sum() / denom
|
||||||
|
|
||||||
|
def train_model():
|
||||||
|
train_loader, val_loader, _, _, _ = get_dataloaders()
|
||||||
|
|
||||||
|
# 初始化前向模型 (1 -> 5)
|
||||||
|
model = BuildingTCN(input_size=1, output_size=5, num_channels=CHANNELS,
|
||||||
|
kernel_size=KERNEL_SIZE, dropout=DROPOUT).to(DEVICE)
|
||||||
|
|
||||||
|
optimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)
|
||||||
|
scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=8, factor=0.5)
|
||||||
|
|
||||||
|
best_val_loss = float('inf')
|
||||||
|
no_improve_epochs = 0
|
||||||
|
early_stop_patience = EARLY_STOP_PATIENCE
|
||||||
|
|
||||||
|
for epoch in range(EPOCHS):
|
||||||
|
model.train()
|
||||||
|
train_loss = 0.0
|
||||||
|
train_time_loss = 0.0
|
||||||
|
train_spec_loss = 0.0
|
||||||
|
train_corr_loss = 0.0
|
||||||
|
train_mse_loss = 0.0
|
||||||
|
train_std_loss = 0.0
|
||||||
|
train_peak_loss = 0.0
|
||||||
|
train_batches = 0
|
||||||
|
|
||||||
|
for inputs, targets, masks in train_loader:
|
||||||
|
inputs = inputs.to(DEVICE)
|
||||||
|
targets = targets.to(DEVICE)
|
||||||
|
masks = masks.to(DEVICE)
|
||||||
|
|
||||||
|
optimizer.zero_grad()
|
||||||
|
outputs = model(inputs)
|
||||||
|
time_loss = masked_l1_loss(outputs, targets, masks)
|
||||||
|
mse_loss = masked_mse_loss(outputs, targets, masks)
|
||||||
|
spec_loss = masked_spectral_mag_loss(outputs, targets, masks)
|
||||||
|
corr_loss = masked_corr_loss(outputs, targets, masks, eps=CORR_LOSS_EPS)
|
||||||
|
std_loss = masked_std_loss(outputs, targets, masks)
|
||||||
|
peak_loss = masked_peak_loss(outputs, targets, masks)
|
||||||
|
loss = (
|
||||||
|
time_loss
|
||||||
|
+ MSE_LOSS_WEIGHT * mse_loss
|
||||||
|
+ SPECTRAL_LOSS_WEIGHT * spec_loss
|
||||||
|
+ CORR_LOSS_WEIGHT * corr_loss
|
||||||
|
+ STD_LOSS_WEIGHT * std_loss
|
||||||
|
+ PEAK_LOSS_WEIGHT * peak_loss
|
||||||
|
)
|
||||||
|
|
||||||
|
if not torch.isfinite(loss):
|
||||||
|
print(" [Warn] Non-finite train loss encountered. Skip this batch.")
|
||||||
|
continue
|
||||||
|
|
||||||
|
loss.backward()
|
||||||
|
grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
||||||
|
if not torch.isfinite(grad_norm):
|
||||||
|
print(" [Warn] Non-finite gradient norm encountered. Skip optimizer step for this batch.")
|
||||||
|
optimizer.zero_grad(set_to_none=True)
|
||||||
|
continue
|
||||||
|
|
||||||
|
optimizer.step()
|
||||||
|
|
||||||
|
train_loss += loss.item()
|
||||||
|
train_time_loss += time_loss.item()
|
||||||
|
train_mse_loss += mse_loss.item()
|
||||||
|
train_spec_loss += spec_loss.item()
|
||||||
|
train_corr_loss += corr_loss.item()
|
||||||
|
train_std_loss += std_loss.item()
|
||||||
|
train_peak_loss += peak_loss.item()
|
||||||
|
train_batches += 1
|
||||||
|
|
||||||
|
# 验证
|
||||||
|
model.eval()
|
||||||
|
val_loss = 0.0
|
||||||
|
val_time_loss = 0.0
|
||||||
|
val_spec_loss = 0.0
|
||||||
|
val_corr_loss = 0.0
|
||||||
|
val_mse_loss = 0.0
|
||||||
|
val_std_loss = 0.0
|
||||||
|
val_peak_loss = 0.0
|
||||||
|
val_batches = 0
|
||||||
|
with torch.no_grad():
|
||||||
|
for inputs, targets, masks in val_loader:
|
||||||
|
inputs = inputs.to(DEVICE)
|
||||||
|
targets = targets.to(DEVICE)
|
||||||
|
masks = masks.to(DEVICE)
|
||||||
|
outputs = model(inputs)
|
||||||
|
time_loss = masked_l1_loss(outputs, targets, masks)
|
||||||
|
mse_loss = masked_mse_loss(outputs, targets, masks)
|
||||||
|
spec_loss = masked_spectral_mag_loss(outputs, targets, masks)
|
||||||
|
corr_loss = masked_corr_loss(outputs, targets, masks, eps=CORR_LOSS_EPS)
|
||||||
|
std_loss = masked_std_loss(outputs, targets, masks)
|
||||||
|
peak_loss = masked_peak_loss(outputs, targets, masks)
|
||||||
|
loss = (
|
||||||
|
time_loss
|
||||||
|
+ MSE_LOSS_WEIGHT * mse_loss
|
||||||
|
+ SPECTRAL_LOSS_WEIGHT * spec_loss
|
||||||
|
+ CORR_LOSS_WEIGHT * corr_loss
|
||||||
|
+ STD_LOSS_WEIGHT * std_loss
|
||||||
|
+ PEAK_LOSS_WEIGHT * peak_loss
|
||||||
|
)
|
||||||
|
if not torch.isfinite(loss):
|
||||||
|
continue
|
||||||
|
val_loss += loss.item()
|
||||||
|
val_time_loss += time_loss.item()
|
||||||
|
val_mse_loss += mse_loss.item()
|
||||||
|
val_spec_loss += spec_loss.item()
|
||||||
|
val_corr_loss += corr_loss.item()
|
||||||
|
val_std_loss += std_loss.item()
|
||||||
|
val_peak_loss += peak_loss.item()
|
||||||
|
val_batches += 1
|
||||||
|
|
||||||
|
train_den = max(train_batches, 1)
|
||||||
|
val_den = max(val_batches, 1)
|
||||||
|
train_loss /= train_den
|
||||||
|
train_time_loss /= train_den
|
||||||
|
train_mse_loss /= train_den
|
||||||
|
train_spec_loss /= train_den
|
||||||
|
train_corr_loss /= train_den
|
||||||
|
train_std_loss /= train_den
|
||||||
|
train_peak_loss /= train_den
|
||||||
|
val_loss /= val_den
|
||||||
|
val_time_loss /= val_den
|
||||||
|
val_mse_loss /= val_den
|
||||||
|
val_spec_loss /= val_den
|
||||||
|
val_corr_loss /= val_den
|
||||||
|
val_std_loss /= val_den
|
||||||
|
val_peak_loss /= val_den
|
||||||
|
|
||||||
|
scheduler.step(val_loss)
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"Epoch {epoch+1}/{EPOCHS} | "
|
||||||
|
f"Train Loss: {train_loss:.4f} "
|
||||||
|
f"(L1={train_time_loss:.4f}, MSE={train_mse_loss:.4f}, Spec={train_spec_loss:.4f}, "
|
||||||
|
f"Corr={train_corr_loss:.4f}, Std={train_std_loss:.4f}, Peak={train_peak_loss:.4f}) | "
|
||||||
|
f"Val Loss: {val_loss:.4f} "
|
||||||
|
f"(L1={val_time_loss:.4f}, MSE={val_mse_loss:.4f}, Spec={val_spec_loss:.4f}, "
|
||||||
|
f"Corr={val_corr_loss:.4f}, Std={val_std_loss:.4f}, Peak={val_peak_loss:.4f})"
|
||||||
|
)
|
||||||
|
|
||||||
|
if val_loss < best_val_loss:
|
||||||
|
best_val_loss = val_loss
|
||||||
|
torch.save(model.state_dict(), 'best_model.pth')
|
||||||
|
print(" --> Saved Best Model")
|
||||||
|
no_improve_epochs = 0
|
||||||
|
else:
|
||||||
|
no_improve_epochs += 1
|
||||||
|
if ENABLE_EARLY_STOP and no_improve_epochs >= early_stop_patience:
|
||||||
|
print(f"Early stopping at epoch {epoch+1}")
|
||||||
|
break
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
train_model()
|
||||||
BIN
test_results.png
BIN
test_results.png
Binary file not shown.
|
Before Width: | Height: | Size: 362 KiB |
Reference in New Issue
Block a user