non-intrusive-se / scripts /inspect_ds_coef.py
vvwangvv's picture
Add non-intrusive DS reference implementation: code, data, checkpoints, recipe
ea2f14c verified
Raw
History Blame Contribute Delete
3.94 kB
"""Diagnostic: inspect the DS coefficient S produced by a trained DS module.
If S is ~constant (low std across band & time), the DS has effectively
re-learned the global OA coefficient and cannot beat OA. If S varies across
bands/time, the DS is adaptive and the OA-vs-DS gap is bounded by SE quality.
"""
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.data import NoisyMixDataset, collate_full # noqa: E402
from ds4se.ds4bsrnn import DS4BSRNN # noqa: E402
from ds4se.ds_module import DSModule # noqa: E402
from ds4se.se_io import load_frozen_se # 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 main():
ap = argparse.ArgumentParser()
ap.add_argument("--se", default=os.path.join(ROOT, "ckpt", "se_bsrnn.pt"))
ap.add_argument("--ds", required=True)
ap.add_argument("--hop", type=int, default=128)
ap.add_argument("--n", type=int, default=40)
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"])
ck = torch.load(args.ds, map_location=device)
if ck["coupled"]:
ds = DS4BSRNN(bsrnn_num_channel=cfg["num_channel"], subbands=se.subbands,
num_layer=1, lstm_hidden=16, mode=ck["mode"], warmup_steps=0,
proj_channel=16)
else:
ds = DSModule(n_fft=cfg["n_fft"], num_channel=16, num_layer=1,
lstm_hidden=16, subbands=se.subbands, mode=ck["mode"],
warmup_steps=0)
ds.load_state_dict(ck["model"])
ds.to(device).eval()
loader = DataLoader(NoisyMixDataset(os.path.join(MANIFESTS, "test.json"),
segment_sec=None, snr_range=(-5, 20), seed=13,
return_text=True), batch_size=4, shuffle=False,
num_workers=2, collate_fn=collate_full)
flat = [] # flattened S values for global stats
std_time_acc, std_time_n = 0.0, 0 # std over time, averaged over (utt,band)
std_band_acc, std_band_n = 0.0, 0 # std over bands, averaged over (utt,time)
band_sum, band_cnt = None, 0 # per-band mean accumulation
with torch.no_grad():
seen = 0
for noisy, clean, lengths, text in loader:
noisy = noisy.to(device)
X = stft.stft(noisy)
enh, H = se(X, return_hidden=True)
_, S = ds(X, enh, H) if ck["coupled"] else ds(X, enh) # (B, K, T)
S = S.cpu()
flat.append(S.reshape(-1))
std_time_acc += S.std(dim=2).sum().item()
std_time_n += S.size(0) * S.size(1)
std_band_acc += S.std(dim=1).sum().item()
std_band_n += S.size(0) * S.size(2)
bs = S.mean(dim=2).sum(dim=0) # (K,)
band_sum = bs if band_sum is None else band_sum + bs
band_cnt += S.size(0)
seen += noisy.size(0)
if seen >= args.n:
break
flat = torch.cat(flat)
band_mean = band_sum / band_cnt
print(f"DS={'coupled' if ck['coupled'] else 'decoupled'} mode={ck['mode']} "
f"dev_wer={ck.get('dev_wer'):.2f} step={ck.get('step')}")
print(f"global mean={flat.mean():.4f} std={flat.std():.4f} "
f"min={flat.min():.4f} max={flat.max():.4f}")
print(f"avg std over TIME (temporal adaptivity) = {std_time_acc/std_time_n:.4f}")
print(f"avg std over BANDS (spectral adaptivity) = {std_band_acc/std_band_n:.4f}")
print("per-band mean S (low->high freq):")
print(" " + " ".join(f"{v:.2f}" for v in band_mean.tolist()))
if __name__ == "__main__":
main()