File size: 7,872 Bytes
1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf 8131a83 1ebfbbf | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 | #!/usr/bin/env python3
"""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) # 等价 torch.hann_window(512, periodic=True)
N_FREQ = N_FFT // 2 + 1 # 257
MAX_FRAMES = 626 # feat[1,6,626,257],约 10.0s @16kHz
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) # [T,ch,N]
return np.fft.rfft(frames * WIN, axis=2).transpose(1, 2, 0) # [ch,F,T]
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 # [T,N]
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] # center 填充裁剪(对齐 torch.istft(center=True))
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) # [F,M,T]
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) # [F,1,T]
temp = X_delay / lambdaa
R = temp @ np.conj(X_delay.transpose(0, 2, 1)) # [F,36,36]
P = temp @ np.conj(Xp.transpose(0, 2, 1)) # [F,36,M]
G = np.linalg.solve(R + eps * np.eye(M * Lg), P) # [F,36,M]
Y = Xp - np.conj(G.transpose(0, 2, 1)) @ X_delay
return Y.transpose(1, 0, 2) # [M,F,T]
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() # [F,M,T]
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 # [F,M,T]
if model == "laplace":
r = 2.0 * np.sqrt((Y.real ** 2 + Y.imag ** 2).sum(axis=0)) # [M,T],沿 F 求和
else:
r = (Y.real ** 2 + Y.imag ** 2).sum(axis=0) / F # 沿 F 求和
r[r < eps] = eps
r_inv = 1.0 / r # [M,T]
for s in range(n_src):
V = ((Xp * r_inv[None, s, None, :]) @ np.conj(Xp.transpose(0, 2, 1))) / T # [F,M,M]
WV = W @ V
e_s = np.tile(eyes[:, s], (F, 1))[:, :, None] # [F,M,1]
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 # [F,M,T]
if proj_back:
# ref = X 的第 0 通道 [T,F];对 T(帧)求和 → 每个 (f, 声源) 一个缩放系数
ref = X[:, :, 0] # [T,F]
num = np.sum(np.conj(ref).T[:, None, :] * Y, axis=2) # [F,M]
denom = np.sum(np.abs(Y) ** 2, axis=2) # [F,M]
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) # [T,F,M]
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) # [M,L]
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:
# 补零到 626 帧(零帧对 WPE/IVA 统计无贡献,proj_back 缩放不变性保证结果一致)
x = np.pad(x, ((0, 0), (0, (MAX_FRAMES - T) * HOP)))
T = MAX_FRAMES
spec_orig = stft(x) # [M,F,T]
spec_drb = fd_wpe(spec_orig) # WPE
spec_2ch = auxiva(spec_drb.transpose(2, 1, 0)).transpose(2, 1, 0) # [M,F,T]
# 通道选择:能量较小的通道
spec_norm = np.sqrt((spec_2ch.real ** 2 + spec_2ch.imag ** 2).sum(axis=(1, 2))) # [M]
pred = 1 if spec_norm[0] < spec_norm[1] else 0
spec_selected = spec_2ch[pred] # [F,T]
spec_unselected = spec_2ch[1 - pred]
sel_log = np.log10(np.maximum(np.abs(spec_selected), 1e-12)).T[None] # [1,T,F]
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) # [4,T,F]
feat = np.concatenate([spec, sel_log, un_log], axis=0)[None] # [1,6,T,F]
feat = np.ascontiguousarray(feat, dtype=np.float32)
session = ModelSession(args.model)
mask = session.run_named({"feat": feat})["mask"] # [1,2,T,F]
m = mask[0] # [2,T,F]
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 # [F,T]
out = istft(spec_enh)[:L]
write_wav(args.output, out)
print("saved:", args.output)
if __name__ == "__main__":
main()
|