| |
| """H-GTCRN 板端 wav→wav 降噪 demo(numpy 实现,依赖仅 numpy + pyaxengine)。 |
| |
| CPU 链路(与官方 GTCRN_IVA 一致):STFT → WPE 去混响 → auxIVA 分离 → |
| 通道选择 + 特征构造 → [NPU: GTCRN 核] → 复数掩码应用 → ISTFT。 |
| 仅支持 16kHz PCM16 wav;长度 ≤ 10.0s(626 帧)自动补零,超过报错。 |
| """ |
| import argparse |
| import sys |
| import wave |
| from pathlib import Path |
|
|
| import numpy as np |
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parents[1])) |
| from h_gtcrn_core_sdk import ModelSession |
|
|
| N_FFT = 512 |
| HOP = 256 |
| WIN = np.hanning(N_FFT + 1)[:-1].astype(np.float64) |
| N_FREQ = N_FFT // 2 + 1 |
| MAX_FRAMES = 626 |
| SR = 16000 |
|
|
|
|
| def read_wav(path): |
| """读 PCM16 wav → [ch, L] float64;单声道复制为双声道。""" |
| with wave.open(path, "rb") as w: |
| assert w.getframerate() == SR, f"expected 16000 Hz, got {w.getframerate()}" |
| assert w.getsampwidth() == 2, "only 16-bit PCM wav supported" |
| ch = w.getnchannels() |
| x = np.frombuffer(w.readframes(w.getnframes()), dtype=np.int16).astype(np.float64) / 32768.0 |
| x = x.reshape(-1, ch).T |
| if ch == 1: |
| x = np.stack([x[0], x[0]], axis=0) |
| elif ch > 2: |
| x = x[:2] |
| return np.ascontiguousarray(x) |
|
|
|
|
| def write_wav(path, x, sr=SR): |
| pcm = np.clip(x, -1.0, 1.0) |
| with wave.open(path, "wb") as w: |
| w.setnchannels(1) |
| w.setsampwidth(2) |
| w.setframerate(sr) |
| w.writeframes((pcm * 32767).astype(np.int16).tobytes()) |
|
|
|
|
| def stft(x): |
| """x [ch, L] → [ch, F, T] complex(center=True 反射填充)。""" |
| ch, L = x.shape |
| xp = np.pad(x, ((0, 0), (N_FFT // 2, N_FFT // 2)), mode="reflect") |
| T = (xp.shape[1] - N_FFT) // HOP + 1 |
| frames = np.stack([xp[:, t * HOP:t * HOP + N_FFT] for t in range(T)], axis=0) |
| return np.fft.rfft(frames * WIN, axis=2).transpose(1, 2, 0) |
|
|
|
|
| def istft(spec): |
| """spec [F, T] complex → (T-1)*hop + n_fft 长度信号,裁掉 center 填充。""" |
| F, T = spec.shape |
| frames = np.fft.irfft(spec.T, n=N_FFT, axis=1) * WIN |
| out_len = (T - 1) * HOP + N_FFT |
| out = np.zeros(out_len, dtype=np.float32) |
| wsum = np.zeros(out_len, dtype=np.float32) |
| for t in range(T): |
| s = t * HOP |
| out[s:s + N_FFT] += frames[t] |
| wsum[s:s + N_FFT] += WIN * WIN |
| wsum[wsum < 1e-10] = 1.0 |
| out = out / wsum |
| pad = N_FFT // 2 |
| return out[pad:out_len - pad] |
|
|
|
|
| def fd_wpe(X, rt60=0.3, shift=256, D=2, fs=16000, num_iter=1): |
| """WPE 去混响。X [M,F,T] complex → [M,F,T] complex(与 torch 版逐式对齐)。""" |
| M, F, T = X.shape |
| eps = 1e-3 * np.mean(np.max(np.max(np.abs(X) ** 2, axis=-1), axis=0)) |
| Lg = int(rt60 * fs / shift) |
| Xp = X.transpose(1, 0, 2) |
| X_delay = np.zeros((F, M * Lg, T), dtype=X.dtype) |
| for l in range(Lg): |
| X_delay[:, l * M:(l + 1) * M, D + l:T] = Xp[:, :, 0:T - D - l] |
| Y = Xp.copy() |
| for _ in range(num_iter): |
| lambdaa = np.maximum(np.mean(np.abs(Y) ** 2, axis=-2, keepdims=True), eps) |
| temp = X_delay / lambdaa |
| R = temp @ np.conj(X_delay.transpose(0, 2, 1)) |
| P = temp @ np.conj(Xp.transpose(0, 2, 1)) |
| G = np.linalg.solve(R + eps * np.eye(M * Lg), P) |
| Y = Xp - np.conj(G.transpose(0, 2, 1)) @ X_delay |
| return Y.transpose(1, 0, 2) |
|
|
|
|
| def auxiva(X, n_src=2, n_iter=10, proj_back=True, model="laplace"): |
| """auxIVA 声源分离。X [T,F,M] complex → [T,F,M] complex(与 torch 版逐式对齐)。""" |
| T, F, M = X.shape |
| eps = 1e-10 |
| eyes = np.eye(M, dtype=X.dtype) |
| W = np.zeros((F, M, M), dtype=X.dtype) |
| W[:, :, :] = np.tile(np.eye(n_src, dtype=X.dtype), (F, 1, 1))[:, :n_src, :] |
| Xp = X.transpose(1, 2, 0).copy() |
| Y = np.zeros((F, M, T), dtype=X.dtype) |
| r = np.zeros((M, T), dtype=np.float32) |
| for _ in range(n_iter): |
| Y = W @ Xp |
| if model == "laplace": |
| r = 2.0 * np.sqrt((Y.real ** 2 + Y.imag ** 2).sum(axis=0)) |
| else: |
| r = (Y.real ** 2 + Y.imag ** 2).sum(axis=0) / F |
| r[r < eps] = eps |
| r_inv = 1.0 / r |
| for s in range(n_src): |
| V = ((Xp * r_inv[None, s, None, :]) @ np.conj(Xp.transpose(0, 2, 1))) / T |
| WV = W @ V |
| e_s = np.tile(eyes[:, s], (F, 1))[:, :, None] |
| W[:, s, :] = np.conj(np.linalg.solve(WV + eps * eyes, e_s))[:, :, 0] |
| denom = (W[:, None, s, :] @ V @ np.conj(W[:, s, :, None]))[:, :, 0] |
| W[:, s, :] /= np.sqrt(denom + eps) |
| Y = W @ Xp |
| if proj_back: |
| |
| ref = X[:, :, 0] |
| num = np.sum(np.conj(ref).T[:, None, :] * Y, axis=2) |
| denom = np.sum(np.abs(Y) ** 2, axis=2) |
| c = np.ones((F, M), dtype=X.dtype) |
| valid = denom > 0.0 |
| c[valid] = num[valid] / denom[valid] |
| Y *= np.conj(c)[:, :, None] |
| return Y.transpose(2, 0, 1) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="H-GTCRN AX650 板端 wav→wav 降噪(numpy 版)。") |
| parser.add_argument("--model", default="../models/model.axmodel") |
| parser.add_argument("--input-wav", default="../samples/Samples1_noisy.wav") |
| parser.add_argument("--output", default="../samples/Samples1_board_enhanced.wav") |
| args = parser.parse_args() |
|
|
| x = read_wav(args.input_wav) |
| L = x.shape[1] |
| T = L // HOP + 1 |
| if T > MAX_FRAMES: |
| raise ValueError(f"input too long: {T} frames > {MAX_FRAMES} (约 10.0s @16kHz)") |
| if T < MAX_FRAMES: |
| |
| x = np.pad(x, ((0, 0), (0, (MAX_FRAMES - T) * HOP))) |
| T = MAX_FRAMES |
|
|
| spec_orig = stft(x) |
| spec_drb = fd_wpe(spec_orig) |
| spec_2ch = auxiva(spec_drb.transpose(2, 1, 0)).transpose(2, 1, 0) |
|
|
| |
| spec_norm = np.sqrt((spec_2ch.real ** 2 + spec_2ch.imag ** 2).sum(axis=(1, 2))) |
| pred = 1 if spec_norm[0] < spec_norm[1] else 0 |
| spec_selected = spec_2ch[pred] |
| spec_unselected = spec_2ch[1 - pred] |
|
|
| sel_log = np.log10(np.maximum(np.abs(spec_selected), 1e-12)).T[None] |
| un_log = np.log10(np.maximum(np.abs(spec_unselected), 1e-12)).T[None] |
|
|
| spec = np.stack([spec_orig[0].real, spec_orig[0].imag, |
| spec_orig[1].real, spec_orig[1].imag], axis=0).transpose(0, 2, 1) |
| feat = np.concatenate([spec, sel_log, un_log], axis=0)[None] |
| feat = np.ascontiguousarray(feat, dtype=np.float32) |
| session = ModelSession(args.model) |
| mask = session.run_named({"feat": feat})["mask"] |
| m = mask[0] |
|
|
| s_real = spec[0] * m[0] - spec[1] * m[1] |
| s_imag = spec[1] * m[0] + spec[0] * m[1] |
| spec_enh = (s_real + 1j * s_imag).T |
| out = istft(spec_enh)[:L] |
| write_wav(args.output, out) |
| print("saved:", args.output) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|