modified: src/config.py

modified:   src/train.py
This commit is contained in:
CrbnsCat10n
2026-04-30 16:43:34 +08:00
parent 05101ac7c7
commit c65e47e6b0
3 changed files with 38 additions and 5 deletions

View File

@@ -24,7 +24,9 @@ EPOCHS = 50
WEIGHT_DECAY = 1e-3
ENABLE_EARLY_STOP = False
EARLY_STOP_PATIENCE = 15
SPECTRAL_LOSS_WEIGHT = 0.2
SPECTRAL_LOSS_WEIGHT = 0.6
CORR_LOSS_WEIGHT = 0.3
CORR_LOSS_EPS = 1e-8
# Device
import torch

View File

@@ -27,6 +27,29 @@ def masked_spectral_mag_loss(pred, target, mask):
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)
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: