"""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()