"""Evaluate WER on the test set for each method (mirrors the paper's tables): Noisy | SE-enhanced | + OA (tuned) | + DS | + DS4BSRNN OA (observation adding) uses a single global coefficient s_OA tuned on dev: X_hat = s_OA * X + (1 - s_OA) * X_tilde Usage: python scripts/eval_wer.py --se ckpt/se_bsrnn.pt \ --ds ckpt/ds.pt --ds4 ckpt/ds4bsrnn.pt """ import argparse import os import sys import torch from torch.utils.data import DataLoader sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from ds4se.asr_whisper import WhisperASR # noqa: E402 from ds4se.data import NoisyMixDataset, collate_full # noqa: E402 from ds4se.ds4bsrnn import DS4BSRNN # noqa: E402 from ds4se.ds_module import DSModule, ds_interpolate # noqa: E402 from ds4se.se_io import load_frozen_se # noqa: E402 from ds4se.stft import STFT # noqa: E402 from ds4se.wer import wer # noqa: E402 ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) MANIFESTS = os.path.join(ROOT, "data", "manifests") @torch.no_grad() def run_method(asr, refine_fn, loader, device, max_items=None): refs, hyps = [], [] for noisy, clean, lengths, text in loader: noisy = noisy.to(device) wav = refine_fn(noisy) # (B, samples) for i in range(noisy.size(0)): w = wav[i : i + 1, : int(lengths[i])] hyps.append(asr.transcribe(w)[0]) refs.append(text[i]) if max_items and len(refs) >= max_items: break return wer(refs, hyps) def make_loader(split, batch_size, snr_range=(-5, 20)): ds = NoisyMixDataset(os.path.join(MANIFESTS, f"{split}.json"), segment_sec=None, snr_range=snr_range, seed=7 if split == "dev" else 13, return_text=True) return DataLoader(ds, batch_size=batch_size, shuffle=False, num_workers=2, collate_fn=collate_full) def main(): ap = argparse.ArgumentParser() ap.add_argument("--se", type=str, default=os.path.join(ROOT, "ckpt", "se_bsrnn.pt")) ap.add_argument("--ds", type=str, default=None) ap.add_argument("--ds4", type=str, default=None) ap.add_argument("--whisper", type=str, default="openai/whisper-base.en") ap.add_argument("--batch_size", type=int, default=4) ap.add_argument("--hop", type=int, default=128) ap.add_argument("--max_items", type=int, default=None) ap.add_argument("--snr_lo", type=float, default=-5.0) ap.add_argument("--snr_hi", type=float, default=20.0) args = ap.parse_args() device = "cuda" if torch.cuda.is_available() else "cpu" se, cfg = load_frozen_se(args.se, device) stft = STFT(n_fft=cfg["n_fft"], hop_length=args.hop, win_length=cfg["n_fft"]) asr = WhisperASR(model_name=args.whisper, device=device) snr_range = (args.snr_lo, args.snr_hi) print(f"snr_range = {snr_range}", flush=True) dev_loader = make_loader("dev", args.batch_size, snr_range) test_loader = make_loader("test", args.batch_size, snr_range) def f_noisy(noisy): return noisy def f_enh(noisy): X = stft.stft(noisy) Y = se(X) return stft.istft(Y, length=noisy.size(-1)) def f_oa(s_oa): def fn(noisy): X = stft.stft(noisy) Y = se(X) Xc = torch.view_as_complex(X.contiguous()) Yc = torch.view_as_complex(Y.contiguous()) X_hat = s_oa * Xc + (1.0 - s_oa) * Yc return stft.istft(X_hat, length=noisy.size(-1)) return fn results = {} print("== evaluating (test set) ==", flush=True) results["Noisy"] = run_method(asr, f_noisy, test_loader, device, args.max_items) print(f"Noisy WER = {results['Noisy']:.2f}%", flush=True) results["SE-enhanced"] = run_method(asr, f_enh, test_loader, device, args.max_items) print(f"SE-enhanced WER = {results['SE-enhanced']:.2f}%", flush=True) # tune OA coefficient on dev best_oa, best_oa_wer = 0.0, 1e9 print("-- tuning OA on dev --", flush=True) for s in [0.2, 0.4, 0.6, 0.8]: w = run_method(asr, f_oa(s), dev_loader, device) # full dev set print(f" s_OA={s:.1f} dev WER={w:.2f}%", flush=True) if w < best_oa_wer: best_oa_wer, best_oa = w, s results[f"+OA(s={best_oa:.1f})"] = run_method(asr, f_oa(best_oa), test_loader, device, args.max_items) print(f"+OA WER = {results[f'+OA(s={best_oa:.1f})']:.2f}% " f"(s_OA={best_oa})", flush=True) def load_ds(path): ck = torch.load(path, map_location=device) coupled = ck["coupled"] mode = ck["mode"] if coupled: m = DS4BSRNN(bsrnn_num_channel=cfg["num_channel"], subbands=se.subbands, num_layer=1, lstm_hidden=16, mode=mode, warmup_steps=0, proj_channel=16) else: m = DSModule(n_fft=cfg["n_fft"], num_channel=16, num_layer=1, lstm_hidden=16, subbands=se.subbands, mode=mode, warmup_steps=0) m.load_state_dict(ck["model"]) m.to(device).eval() return m, coupled if args.ds: ds, coupled = load_ds(args.ds) def f_ds(noisy): X = stft.stft(noisy) enh, H = se(X, return_hidden=True) X_hat, _ = ds(X, enh, H) if coupled else ds(X, enh) return stft.istft(X_hat, length=noisy.size(-1)) results["+DS"] = run_method(asr, f_ds, test_loader, device, args.max_items) print(f"+DS WER = {results['+DS']:.2f}%", flush=True) if args.ds4: ds4, coupled4 = load_ds(args.ds4) def f_ds4(noisy): X = stft.stft(noisy) enh, H = se(X, return_hidden=True) X_hat, _ = ds4(X, enh, H) if coupled4 else ds4(X, enh) return stft.istft(X_hat, length=noisy.size(-1)) results["+DS4BSRNN"] = run_method(asr, f_ds4, test_loader, device, args.max_items) print(f"+DS4BSRNN WER = {results['+DS4BSRNN']:.2f}%", flush=True) print("\n==== TEST WER SUMMARY ====") for k, v in results.items(): print(f" {k:16s} {v:6.2f}%") if __name__ == "__main__": main()