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