non-intrusive-se / scripts /train_se.py
vvwangvv's picture
Add non-intrusive DS reference implementation: code, data, checkpoints, recipe
ea2f14c verified
Raw
History Blame Contribute Delete
4.87 kB
"""Train a BSRNN speech-enhancement front-end on simulated noisy/clean data.
This is the (frozen) SE model the DS module later corrects. It is trained with
an SI-SDR waveform loss. Not aiming for SOTA -- just a real SE front-end that
denoises while introducing the kind of distortion DS is meant to suppress.
Usage:
python scripts/train_se.py --epochs 30 --out ckpt/se_bsrnn.pt
"""
import argparse
import os
import sys
import time
import torch
from torch.utils.data import DataLoader
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from ds4se.bsrnn import BSRNN # noqa: E402
from ds4se.data import NoisyMixDataset, collate_segments # noqa: E402
from ds4se.losses import si_sdr, si_sdr_loss # noqa: E402
from ds4se.stft import STFT # noqa: E402
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
MANIFESTS = os.path.join(ROOT, "data", "manifests")
def enhance(model, stft, noisy, length):
X = stft.stft(noisy)
Y = model(X)
return stft.istft(Y, length=length)
@torch.no_grad()
def evaluate(model, stft, loader, device, max_batches=20):
model.eval()
vals = []
for i, (noisy, clean) in enumerate(loader):
if i >= max_batches:
break
noisy, clean = noisy.to(device), clean.to(device)
est = enhance(model, stft, noisy, clean.size(-1))
vals.append(si_sdr(est[..., : clean.size(-1)], clean).mean().item())
return sum(vals) / max(len(vals), 1)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--epochs", type=int, default=30)
ap.add_argument("--batch_size", type=int, default=16)
ap.add_argument("--segment_sec", type=float, default=4.0)
ap.add_argument("--lr", type=float, default=1e-3)
ap.add_argument("--num_channel", type=int, default=64)
ap.add_argument("--num_layer", type=int, default=6)
ap.add_argument("--n_fft", type=int, default=512)
ap.add_argument("--hop", type=int, default=128)
ap.add_argument("--num_workers", type=int, default=6)
ap.add_argument("--out", type=str, default=os.path.join(ROOT, "ckpt", "se_bsrnn.pt"))
args = ap.parse_args()
device = "cuda" if torch.cuda.is_available() else "cpu"
os.makedirs(os.path.dirname(args.out), exist_ok=True)
train_ds = NoisyMixDataset(os.path.join(MANIFESTS, "train.json"),
segment_sec=args.segment_sec, snr_range=(-5, 20))
dev_ds = NoisyMixDataset(os.path.join(MANIFESTS, "dev.json"),
segment_sec=args.segment_sec, snr_range=(-5, 20), seed=7)
train_loader = DataLoader(train_ds, batch_size=args.batch_size, shuffle=True,
num_workers=args.num_workers, collate_fn=collate_segments,
drop_last=True, persistent_workers=args.num_workers > 0)
dev_loader = DataLoader(dev_ds, batch_size=args.batch_size, shuffle=False,
num_workers=2, collate_fn=collate_segments)
model = BSRNN(n_fft=args.n_fft, num_channel=args.num_channel,
num_layer=args.num_layer).to(device)
stft = STFT(n_fft=args.n_fft, hop_length=args.hop, win_length=args.n_fft)
opt = torch.optim.Adam(model.parameters(), lr=args.lr, weight_decay=1e-5)
sched = torch.optim.lr_scheduler.StepLR(opt, step_size=1, gamma=0.95)
n_params = sum(p.numel() for p in model.parameters()) / 1e6
print(f"BSRNN SE: {n_params:.2f}M params, {model.num_bands} bands, "
f"device={device}")
print(f"train={len(train_ds)} utts, dev={len(dev_ds)} utts")
best = -1e9
for epoch in range(args.epochs):
model.train()
t0, run = time.time(), 0.0
for step, (noisy, clean) in enumerate(train_loader):
noisy, clean = noisy.to(device), clean.to(device)
est = enhance(model, stft, noisy, clean.size(-1))
loss = si_sdr_loss(est, clean)
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)
opt.step()
run += loss.item()
sched.step()
dev_sisdr = evaluate(model, stft, dev_loader, device)
print(f"epoch {epoch+1:02d}/{args.epochs} "
f"train_loss={run/(step+1):.3f} dev_SI-SDR={dev_sisdr:.2f}dB "
f"({time.time()-t0:.0f}s)", flush=True)
if dev_sisdr > best:
best = dev_sisdr
torch.save({"model": model.state_dict(),
"config": {"n_fft": args.n_fft, "hop": args.hop,
"num_channel": args.num_channel,
"num_layer": args.num_layer},
"dev_sisdr": dev_sisdr, "epoch": epoch + 1}, args.out)
print(f"done. best dev SI-SDR={best:.2f} dB. ckpt: {args.out}")
if __name__ == "__main__":
main()