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