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:
0
src_old/__init__.py
Normal file
0
src_old/__init__.py
Normal file
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()
|
||||
Reference in New Issue
Block a user