Compare commits
2 Commits
main
...
c65e47e6b0
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c65e47e6b0 | ||
|
|
05101ac7c7 |
Binary file not shown.
@@ -24,6 +24,9 @@ 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.6
|
||||||
|
CORR_LOSS_WEIGHT = 0.3
|
||||||
|
CORR_LOSS_EPS = 1e-8
|
||||||
|
|
||||||
# Device
|
# Device
|
||||||
import torch
|
import torch
|
||||||
|
|||||||
76
src/train.py
76
src/train.py
@@ -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,47 @@ 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 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():
|
def train_model():
|
||||||
train_loader, val_loader, _, _, _ = get_dataloaders()
|
train_loader, val_loader, _, _, _ = get_dataloaders()
|
||||||
|
|
||||||
@@ -27,6 +67,9 @@ 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
|
||||||
|
train_corr_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 +78,57 @@ 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)
|
||||||
|
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()
|
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()
|
||||||
|
train_corr_loss += corr_loss.item()
|
||||||
|
|
||||||
# 验证
|
# 验证
|
||||||
model.eval()
|
model.eval()
|
||||||
val_loss = 0.0
|
val_loss = 0.0
|
||||||
|
val_time_loss = 0.0
|
||||||
|
val_spec_loss = 0.0
|
||||||
|
val_corr_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)
|
||||||
|
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_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_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_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)
|
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}, 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:
|
if val_loss < best_val_loss:
|
||||||
best_val_loss = val_loss
|
best_val_loss = val_loss
|
||||||
|
|||||||
Reference in New Issue
Block a user