non-intrusive-se / scripts /eval_wer.py
vvwangvv's picture
Add non-intrusive DS reference implementation: code, data, checkpoints, recipe
ea2f14c verified
Raw
History Blame Contribute Delete
6.32 kB
"""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()