From 05101ac7c757c422e216e91eb4b3a10651e7bf28 Mon Sep 17 00:00:00 2001 From: CrbnsCat10n Date: Thu, 30 Apr 2026 16:10:01 +0800 Subject: [PATCH] modified: src/config.py modified: src/train.py --- src/__pycache__/config.cpython-314.pyc | Bin 1047 -> 1079 bytes src/config.py | 1 + src/train.py | 45 ++++++++++++++++++++++--- 3 files changed, 42 insertions(+), 4 deletions(-) diff --git a/src/__pycache__/config.cpython-314.pyc b/src/__pycache__/config.cpython-314.pyc index 505fb57ca6e43f72c320da467b1696e32bb60288..4ec07c536b5e22804e4b2f232bef3d3a86daa6b2 100644 GIT binary patch delta 193 zcmbQvv7Lign~#@^0SFvU=aDt#=yvXLs(+E)kLf7 c!de%FwXO^6ToTsl5Nl%FpxVS%qz%*o0CkKo8UO$Q delta 146 zcmdnaF`a`~n~#@^0SIP2_>?(mBCjN)(?s>R3ULfUjGm$;qOqa@Iv`n85THNVf>E1O zB8X9$A&9p`VsZ}SUPj5uT1=lN8!_uL#!OCN7U2TvV+7*j)syR(`&2)(F);Gp5SEy3 gHPPz2u+~Lkt?R-%mxOgX#G2SPs5Y?`X#v#&00OolZU6uP diff --git a/src/config.py b/src/config.py index 1b610a7..1eb5299 100644 --- a/src/config.py +++ b/src/config.py @@ -24,6 +24,7 @@ EPOCHS = 50 WEIGHT_DECAY = 1e-3 ENABLE_EARLY_STOP = False EARLY_STOP_PATIENCE = 15 +SPECTRAL_LOSS_WEIGHT = 0.2 # Device import torch diff --git a/src/train.py b/src/train.py index e71a811..81f4839 100644 --- a/src/train.py +++ b/src/train.py @@ -1,5 +1,4 @@ import torch -import torch.nn as nn import torch.optim as optim from config import * from dataset import get_dataloaders @@ -10,6 +9,24 @@ def masked_l1_loss(pred, target, mask): denom = mask.sum().clamp(min=1.0) return diff.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) + + # 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 train_model(): train_loader, val_loader, _, _, _ = get_dataloaders() @@ -27,6 +44,8 @@ def train_model(): for epoch in range(EPOCHS): model.train() train_loss = 0.0 + train_time_loss = 0.0 + train_spec_loss = 0.0 for inputs, targets, masks in train_loader: inputs = inputs.to(DEVICE) @@ -35,32 +54,50 @@ def train_model(): optimizer.zero_grad() outputs = model(inputs) - loss = masked_l1_loss(outputs, targets, masks) + 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 loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() train_loss += loss.item() + train_time_loss += time_loss.item() + train_spec_loss += spec_loss.item() # 验证 model.eval() val_loss = 0.0 + val_time_loss = 0.0 + val_spec_loss = 0.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) - loss = masked_l1_loss(outputs, targets, masks) + 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 val_loss += loss.item() + val_time_loss += time_loss.item() + val_spec_loss += spec_loss.item() train_loss /= len(train_loader) + train_time_loss /= len(train_loader) + train_spec_loss /= len(train_loader) val_loss /= len(val_loader) + val_time_loss /= len(val_loader) + val_spec_loss /= len(val_loader) scheduler.step(val_loss) - print(f"Epoch {epoch+1}/{EPOCHS} | Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}") + 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})" + ) if val_loss < best_val_loss: best_val_loss = val_loss