From c65e47e6b0b270c858809a0db9c96efc3054637c Mon Sep 17 00:00:00 2001 From: CrbnsCat10n Date: Thu, 30 Apr 2026 16:43:34 +0800 Subject: [PATCH] modified: src/config.py modified: src/train.py --- src/__pycache__/config.cpython-314.pyc | Bin 1079 -> 1158 bytes src/config.py | 4 ++- src/train.py | 39 ++++++++++++++++++++++--- 3 files changed, 38 insertions(+), 5 deletions(-) diff --git a/src/__pycache__/config.cpython-314.pyc b/src/__pycache__/config.cpython-314.pyc index 4ec07c536b5e22804e4b2f232bef3d3a86daa6b2..26527bab509c9bd47d427a4dd4b03dff341786e6 100644 GIT binary patch delta 257 zcmdna(Zh#l*p=AmHpD6cq2{9~>MX?&|685ps(cA>tYk zJlT-hnM)sNFC!2a@0nc3JjLLKu-F9#fg1uM4ZNS(7#MkP2un=2nrL-hSnHy&)^%Z> TOTs!GVohutRGZj}bb;Cc1N%S$ delta 148 zcmZqU+|I$P&Bx2d00fS_pEK7^ z(Kv=6Mo+O4u~@MHJ&-O`5MVgjfl*VCr$jP{QJEo#7eZ?aOs--Kn*5wGl~HQ5FB1pj d 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) + + corr = cov / (torch.sqrt(pred_var * target_var) + eps) + 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 train_model(): train_loader, val_loader, _, _, _ = get_dataloaders() @@ -46,6 +69,7 @@ def train_model(): train_loss = 0.0 train_time_loss = 0.0 train_spec_loss = 0.0 + train_corr_loss = 0.0 for inputs, targets, masks in train_loader: inputs = inputs.to(DEVICE) @@ -56,7 +80,8 @@ def train_model(): outputs = model(inputs) time_loss = masked_l1_loss(outputs, targets, masks) spec_loss = masked_spectral_mag_loss(outputs, targets, masks) - loss = time_loss + SPECTRAL_LOSS_WEIGHT * spec_loss + 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) @@ -65,12 +90,14 @@ def train_model(): train_loss += loss.item() train_time_loss += time_loss.item() train_spec_loss += spec_loss.item() + train_corr_loss += corr_loss.item() # 验证 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) @@ -79,24 +106,28 @@ def train_model(): outputs = model(inputs) time_loss = masked_l1_loss(outputs, targets, masks) spec_loss = masked_spectral_mag_loss(outputs, targets, masks) - loss = time_loss + SPECTRAL_LOSS_WEIGHT * spec_loss + 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) train_time_loss /= len(train_loader) train_spec_loss /= len(train_loader) + train_corr_loss /= len(train_loader) val_loss /= len(val_loader) val_time_loss /= len(val_loader) val_spec_loss /= len(val_loader) + val_corr_loss /= len(val_loader) scheduler.step(val_loss) print( f"Epoch {epoch+1}/{EPOCHS} | " - f"Train Loss: {train_loss:.4f} (L1={train_time_loss:.4f}, Spec={train_spec_loss:.4f}) | " - f"Val Loss: {val_loss:.4f} (L1={val_time_loss:.4f}, Spec={val_spec_loss:.4f})" + f"Train Loss: {train_loss:.4f} (L1={train_time_loss:.4f}, Spec={train_spec_loss:.4f}, Corr={train_corr_loss:.4f}) | " + f"Val Loss: {val_loss:.4f} (L1={val_time_loss:.4f}, Spec={val_spec_loss:.4f}, Corr={val_corr_loss:.4f})" ) if val_loss < best_val_loss: