| """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 |
| from ds4se.data import NoisyMixDataset, collate_full |
| from ds4se.ds4bsrnn import DS4BSRNN |
| from ds4se.ds_module import DSModule, ds_interpolate |
| from ds4se.se_io import load_frozen_se |
| from ds4se.stft import STFT |
| from ds4se.wer import wer |
|
|
| 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) |
| 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) |
|
|
| |
| 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) |
| 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() |
|
|