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

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()