modified: src/config.py

modified:   src/train.py
This commit is contained in:
CrbnsCat10n
2026-04-30 16:10:01 +08:00
parent c82f61a4b3
commit 05101ac7c7
3 changed files with 42 additions and 4 deletions

View File

@@ -24,6 +24,7 @@ EPOCHS = 50
WEIGHT_DECAY = 1e-3 WEIGHT_DECAY = 1e-3
ENABLE_EARLY_STOP = False ENABLE_EARLY_STOP = False
EARLY_STOP_PATIENCE = 15 EARLY_STOP_PATIENCE = 15
SPECTRAL_LOSS_WEIGHT = 0.2
# Device # Device
import torch import torch

View File

@@ -1,5 +1,4 @@
import torch import torch
import torch.nn as nn
import torch.optim as optim import torch.optim as optim
from config import * from config import *
from dataset import get_dataloaders from dataset import get_dataloaders
@@ -10,6 +9,24 @@ def masked_l1_loss(pred, target, mask):
denom = mask.sum().clamp(min=1.0) denom = mask.sum().clamp(min=1.0)
return diff.sum() / denom 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(): def train_model():
train_loader, val_loader, _, _, _ = get_dataloaders() train_loader, val_loader, _, _, _ = get_dataloaders()
@@ -27,6 +44,8 @@ def train_model():
for epoch in range(EPOCHS): for epoch in range(EPOCHS):
model.train() model.train()
train_loss = 0.0 train_loss = 0.0
train_time_loss = 0.0
train_spec_loss = 0.0
for inputs, targets, masks in train_loader: for inputs, targets, masks in train_loader:
inputs = inputs.to(DEVICE) inputs = inputs.to(DEVICE)
@@ -35,32 +54,50 @@ def train_model():
optimizer.zero_grad() optimizer.zero_grad()
outputs = model(inputs) 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() loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step() optimizer.step()
train_loss += loss.item() train_loss += loss.item()
train_time_loss += time_loss.item()
train_spec_loss += spec_loss.item()
# 验证 # 验证
model.eval() model.eval()
val_loss = 0.0 val_loss = 0.0
val_time_loss = 0.0
val_spec_loss = 0.0
with torch.no_grad(): with torch.no_grad():
for inputs, targets, masks in val_loader: for inputs, targets, masks in val_loader:
inputs = inputs.to(DEVICE) inputs = inputs.to(DEVICE)
targets = targets.to(DEVICE) targets = targets.to(DEVICE)
masks = masks.to(DEVICE) masks = masks.to(DEVICE)
outputs = model(inputs) 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_loss += loss.item()
val_time_loss += time_loss.item()
val_spec_loss += spec_loss.item()
train_loss /= len(train_loader) train_loss /= len(train_loader)
train_time_loss /= len(train_loader)
train_spec_loss /= len(train_loader)
val_loss /= len(val_loader) val_loss /= len(val_loader)
val_time_loss /= len(val_loader)
val_spec_loss /= len(val_loader)
scheduler.step(val_loss) 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: if val_loss < best_val_loss:
best_val_loss = val_loss best_val_loss = val_loss