| """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 |
| from ds4se.ds4bsrnn import DS4BSRNN |
| from ds4se.ds_module import DSModule |
| from ds4se.se_io import load_frozen_se |
| from ds4se.stft import STFT |
|
|
| 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 = [] |
| std_time_acc, std_time_n = 0.0, 0 |
| std_band_acc, std_band_n = 0.0, 0 |
| band_sum, band_cnt = None, 0 |
| 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) |
| 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) |
| 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() |
|
|