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:
2026-05-04 18:18:44 +08:00
parent c65e47e6b0
commit dcc023cc04
34 changed files with 2177 additions and 436 deletions

0
src_old/__init__.py Normal file
View File

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

36
src_old/config.py Normal file
View 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
View 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
View 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
View 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
View 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()