| """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 |
| from ds4se.data import NoisyMixDataset, collate_segments |
| from ds4se.losses import si_sdr, si_sdr_loss |
| from ds4se.stft import STFT |
|
|
| 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() |
|
|