最简单DeepAEC Demo -- ICASSP 2022 AEC Challenge Baseline 算法解读
在这篇博客中,我们将深入探讨 ICASSP 2022 会议上提出的 AEC 基线模型,并提供pytorch实现。该模型旨在通过深度学习技术改善音频信号的质量,特别是在回声消除和噪声抑制方面。
一、BASELINE代码解读
1.1 模型解析
该模型使用 ONNX格式进行推理,输入为麦克风(mic)和远端(far-end)音频信号,得到回声、噪声消除后的语音(voice)。
1.1.1 模型输入
-
input:FLOAT[0, 0, 322]- 这是模型的主要输入,形状为
[0, 0, 322],由于动态输入的特性,输入序列长度不定长。 322是特征维度,表示每一帧的输入特征数量,相当于 mic 信号和far-end信号在时域进行频谱并联。
- 这是模型的主要输入,形状为
-
h01:FLOAT[1, 1, 322]- 这是第一个 GRU 层的初始隐藏状态,形状为
[1, 1, 322],表示有一个批次的隐藏状态,特征维度为322。
- 这是第一个 GRU 层的初始隐藏状态,形状为
-
h02:FLOAT[1, 1, 322]- 这是第二个 GRU 层的初始隐藏状态,形状同样为
[1, 1, 322]。
- 这是第二个 GRU 层的初始隐藏状态,形状同样为
1.1.2 模型输出
-
output:FLOAT[0, 0, 161]- 这是模型的主要输出,形状为
[0, 0, 161],表示输出序列的长度和特征维度。161是每一帧用于重建语音信号的频点数量。
- 这是模型的主要输出,形状为
-
hn1:FLOAT[1, 0, 322]- 这是第一个 GRU 层的最终隐藏状态,形状为
[1, 0, 322],表示在处理完输入后,隐藏状态的维度。
- 这是第一个 GRU 层的最终隐藏状态,形状为
-
hn2:FLOAT[1, 0, 322]- 这是第二个 GRU 层的最终隐藏状态,形状同样为
[1, 0, 322]。
- 这是第二个 GRU 层的最终隐藏状态,形状同样为
1.1.3 模型节点
Model Inputs:
input:
Shape: (1, 1, 322), Type: float32
h01:
Shape: (1, 1, 322), Type: float32
h02:
Shape: (1, 1, 322), Type: float32
After GRU_0 and Squeeze_1:
Shape: (1, 322), Type: float32
After Add_2:
Shape: (1, 1, 322), Type: float32
After GRU_3 and Squeeze_4:
Shape: (1, 322), Type: float32
After MatMul_5:
Shape: (1, 161), Type: float32
After Add_6:
Shape: (1, 161), Type: float32
After Sigmoid_7:
Shape: (1, 161), Type: float32
After Clip_8:
Shape: (1, 161), Type: float32
Final Output:
Shape: (1, 161), Type: float32
在分析模型的每一层时,我们使用以下输入数据进行推理:
- 输入数据:
input: 形状为(1, 1, 322)(表示批次大小为 1,序列长度为 1,特征维度为 322)h01: 形状为(1, 1, 322)(第一个 GRU 的初始隐藏状态)h02: 形状为(1, 1, 322)(第二个 GRU 的初始隐藏状态)
模型输入
Model Inputs:
input:
Shape: (1, 1, 322), Type: float32
h01:
Shape: (1, 1, 322), Type: float32
h02:
Shape: (1, 1, 322), Type: float32
-
GRU 层 (GRU_0)
- 输入:
input:(1, 1, 322)h01:(1, 1, 322)
- 输出:
hn1:(1, 1, 322)(最终隐藏状态)output:(1, 1, 322)(GRU 的输出)- 节点标识符:
34(GRU_0 的输出张量名称)
- 输入:
-
Squeeze 层 (Squeeze_1)
- 输入:
34(GRU_0 的输出,形状为(1, 1, 322))
- 输出:
36:(1, 322)(去掉了序列长度为 1 的维度)- 节点标识符:
36(Squeeze_1 的输出张量名称)
- 输入:
-
Add 层 (Add_2)
- 输入:
input:(1, 1, 322)36:(1, 322)
- 输出:
37:(1, 1, 322)(加法的结果)- 节点标识符:
37(Add_2 的输出张量名称)
- 输入:
-
GRU 层 (GRU_3)
- 输入:
37:(1, 1, 322)h02:(1, 1, 322)
- 输出:
hn2:(1, 1, 322)(最终隐藏状态)output:(1, 1, 322)(GRU 的输出)- 节点标识符:
59(GRU_3 的输出张量名称)
- 输入:
-
Squeeze 层 (Squeeze_4)
- 输入:
59(GRU_3 的输出,形状为(1, 1, 322))
- 输出:
61:(1, 322)(去掉了序列长度为 1 的维度)- 节点标识符:
61(Squeeze_4 的输出张量名称)
- 输入:
-
MatMul 层 (MatMul_5)
- 输入:
61:(1, 322)- 假设权重矩阵形状为
(322, 161)
- 输出:
63:(1, 161)(矩阵乘法的结果)- 节点标识符:
63(MatMul_5 的输出张量名称)
- 输入:
-
Add 层 (Add_6)
- 输入:
63:(1, 161)- 假设偏置形状为
(161)
- 输出:
64:(1, 161)(加法的结果)- 节点标识符:
64(Add_6 的输出张量名称)
- 输入:
-
Sigmoid 层 (Sigmoid_7)
- 输入:
64:(1, 161)
- 输出:
65:(1, 161)(经过 Sigmoid 激活函数后的结果)- 节点标识符:
65(Sigmoid_7 的输出张量名称)
- 输入:
-
Clip 层 (Clip_8)
- 输入:
65:(1, 161)
- 输出:
output:(1, 161)(经过 Clip 操作后的结果)- 节点标识符:
output(Clip_8 的输出张量名称)
- 输入:
最终输出
Final Output:
Shape: (1, 161), Type: float32
1.2 代码解析
1.2.1 初始化
模型的初始化方法 __init__ 接受多个参数,包括模型路径、窗口长度、跳跃因子、DFT 大小、隐藏层大小和采样率。以下是初始化方法的关键部分:
def __init__(self, model_path, window_length, hop_fraction,
dft_size, hidden_size, sampling_rate=16_000):
self.hop_fraction = hop_fraction
self.dft_size = dft_size
self.hidden_size = hidden_size
self.sampling_rate = sampling_rate
self.frame_size = int(window_length * sampling_rate)
self.hop_size = int(window_length * sampling_rate * hop_fraction)
self.window = np.sqrt(np.hanning(int(window_length * sampling_rate) + 1)[:-1]).astype(np.float32)
self.model = onnxruntime.InferenceSession(model_path)
- 窗口长度:定义了每个帧的长度。
- 跳跃因子:决定了帧之间的重叠程度。
- DFT 大小:用于快速傅里叶变换(FFT)的大小。
- 隐藏层大小:模型内部的隐藏状态维度。
- 采样率:音频信号的采样率。
1.2.2 特征计算
calc_features 方法用于计算输入音频信号的特征。它将麦克风和远端信号的幅度谱转换为对数功率谱,并进行归一化处理:
def calc_features(self, xmag_mic, xmag_far):
feat_mic = self.logpow(xmag_mic)
feat_far = self.logpow(xmag_far)
feat = np.concatenate([feat_mic, feat_far])
feat /= 20.
feat = feat[np.newaxis, np.newaxis, :]
feat = feat.astype(np.float32)
return feat
- 对数功率谱:通过
logpow方法计算,避免了数值不稳定性。 - 特征拼接:将麦克风和远端信号的特征拼接在一起,形成输入特征。
1.2.3 增强过程
enhance 方法是模型的核心,负责加载音频文件、进行信号增强和输出重建:
def enhance(self, path_mic, path_far):
# load inputs
x_mic, _ = librosa.load(path_mic, sr=self.sampling_rate)
x_far, _ = librosa.load(path_far, sr=self.sampling_rate)
# cut to equal length
min_len = min(len(x_mic), len(x_far))
x_mic = x_mic[:min_len]
x_far = x_far[:min_len]
# zero pad from left
pad_left, pad_right = self.hop_size, 0
x_mic = np.pad(x_mic, (pad_left, pad_right))
x_far = np.pad(x_far, (pad_left, pad_right))
# init buffers
num_frames = (len(x_mic) - self.frame_size) // self.hop_size + 1
x_back = np.zeros(self.frame_size + (num_frames - 1) * self.hop_size)
h01 = np.zeros((1, 1, self.hidden_size), dtype=np.float32)
h02 = np.zeros((1, 1, self.hidden_size), dtype=np.float32)
# frame-wise inference
for ix_start in range(0, len(x_mic) - self.frame_size, self.hop_size):
ix_end = ix_start + self.frame_size
cspec_mic = np.fft.rfft(x_mic[ix_start:ix_end] * self.window, self.dft_size)
xmag_mic, xphs_mic = self.magphasor(cspec_mic)
cspec_far = np.fft.rfft(x_far[ix_start:ix_end] * self.window)
xmag_far = np.abs(cspec_far)
feat = self.calc_features(xmag_mic, xmag_far)
inputs = {
"input": feat,
"h01": h01,
"h02": h02,
}
mask, h01, h02 = self.model.run(None, inputs)
mask = mask[0, 0]
x_enh = np.fft.irfft(mask * xmag_mic * xphs_mic, self.dft_size) * self.window
x_back[ix_start:ix_end] += x_enh
return x_back[pad_left:]
- 音频加载:使用
librosa加载麦克风和远端音频信号。 - 信号填充:为了处理边界情况,信号在左侧进行零填充。
- 帧处理:通过循环对每一帧进行处理,计算幅度谱和相位谱,并提取特征。
- 模型推理:将特征输入到 ONNX 模型中进行推理,得到增强的掩码。
- 信号重建:使用逆傅里叶变换将增强后的信号重建为时域信号。
1.3 输入输出解析
1.3.1 输入
- 麦克风信号:路径由
path_mic指定,加载后为一维数组。 - 远端信号:路径由
path_far指定,加载后为一维数组。 - 特征输入:经过处理后,特征的形状为
(1, 1, feature_size),其中feature_size是麦克风和远端信号特征的拼接结果。
1.3.2 输出
- 增强后的信号:
enhance方法返回增强后的音频信号,形状与输入信号相同,经过重建后为一维数组。
1.4 使用示例
在主程序中,用户可以通过命令行参数指定模型路径、数据目录和输出目录。以下是主程序的关键部分:
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Baseline model inference")
parser.add_argument("--model_path", "-m", help="ONNX model to use for inference.",
default="dec-baseline-model-icassp2022.onnx")
parser.add_argument("--data_dir", "-d", required=True, help="Directory containing mic and farend files.")
parser.add_argument("--output_dir", "-o", required=True, help="Output directory to save enhanced output files.")
parser.add_argument("--output_sr", "-osr", default=48_000, help="Sample rate for output files.")
args = parser.parse_args()
model = DECModel(
model_path=args.model_path,
window_length=0.02,
hop_fraction=0.5,
dft_size=320,
hidden_size=322,
sampling_rate=model_sampling_rate)
mic_paths = glob.glob(os.path.join(args.data_dir, "*_mic.wav"))
for mic_path in tqdm(mic_paths):
basename = os.path.basename(mic_path)
farend_path = mic_path.replace("_mic.wav", "_lpb.wav")
if not os.path.exists(farend_path):
print("Farend file not found, skipping:", farend_path)
continue
out_path = os.path.join(args.output_dir, basename)
x_enhanced = model.enhance(mic_path, farend_path)
x_enhanced = librosa.resample(x_enhanced, orig_sr=model_sampling_rate, target_sr=args.output_sr)
sf.write(out_path, x_enhanced, args.output_sr)
1.4.1 命令行参数
--model_path:指定 ONNX 模型文件路径。--data_dir:包含麦克风和远端音频文件的目录。--output_dir:保存增强后音频文件的输出目录。--output_sr:输出音频文件的采样率。
1.4.2 处理流程
- 加载模型和音频文件。
- 对每个音频文件进行增强处理。
- 将增强后的音频文件保存到指定目录。
完整代码
import argparse
import glob
import os
import librosa
import numpy as np
import onnxruntime
import soundfile as sf
from tqdm import tqdm
class DECModel:
def __init__(self, model_path, window_length, hop_fraction,
dft_size, hidden_size, sampling_rate=16_000):
self.hop_fraction = hop_fraction
self.dft_size = dft_size
self.hidden_size = hidden_size
self.sampling_rate = sampling_rate
self.frame_size = int(window_length * sampling_rate)
self.hop_size = int(window_length * sampling_rate * hop_fraction)
self.window = np.sqrt(np.hanning(int(window_length * sampling_rate) + 1)[:-1]).astype(np.float32)
self.model = onnxruntime.InferenceSession(model_path)
@staticmethod
def logpow(sig):
pspec = np.maximum(sig ** 2, 1e-12)
return np.log10(pspec)
@staticmethod
def magphasor(complexspec):
mspec = np.abs(complexspec)
pspec = np.empty_like(complexspec)
zero_mag = mspec == 0.
pspec[zero_mag] = 1.
pspec[~zero_mag] = complexspec[~zero_mag] / mspec[~zero_mag]
return mspec, pspec
def calc_features(self, xmag_mic, xmag_far):
feat_mic = self.logpow(xmag_mic)
feat_far = self.logpow(xmag_far)
feat = np.concatenate([feat_mic, feat_far])
feat /= 20.
feat = feat[np.newaxis, np.newaxis, :]
feat = feat.astype(np.float32)
return feat
def enhance(self, path_mic, path_far):
# load inputs
x_mic, _ = librosa.load(path_mic, sr=self.sampling_rate)
x_far, _ = librosa.load(path_far, sr=self.sampling_rate)
# cut to equal length
min_len = min(len(x_mic), len(x_far))
x_mic = x_mic[:min_len]
x_far = x_far[:min_len]
# zero pad from left
pad_left, pad_right = self.hop_size, 0
x_mic = np.pad(x_mic, (pad_left, pad_right))
x_far = np.pad(x_far, (pad_left, pad_right))
# init buffers
num_frames = (len(x_mic) - self.frame_size) // self.hop_size + 1
x_back = np.zeros(self.frame_size + (num_frames - 1) * self.hop_size)
h01 = np.zeros((1, 1, self.hidden_size), dtype=np.float32)
h02 = np.zeros((1, 1, self.hidden_size), dtype=np.float32)
# frame-wise inference
for ix_start in range(0, len(x_mic) - self.frame_size, self.hop_size):
ix_end = ix_start + self.frame_size
cspec_mic = np.fft.rfft(x_mic[ix_start:ix_end] * self.window, self.dft_size)
xmag_mic, xphs_mic = self.magphasor(cspec_mic)
cspec_far = np.fft.rfft(x_far[ix_start:ix_end] * self.window)
xmag_far = np.abs(cspec_far)
feat = self.calc_features(xmag_mic, xmag_far)
inputs = {
"input": feat,
"h01": h01,
"h02": h02,
}
mask, h01, h02 = self.model.run(None, inputs)
mask = mask[0, 0]
x_enh = np.fft.irfft(mask * xmag_mic * xphs_mic, self.dft_size) * self.window
x_back[ix_start:ix_end] += x_enh
return x_back[pad_left:]
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Baseline model inference")
parser.add_argument("--model_path", "-m", help="ONNX model to use for inference.",
default="dec-baseline-model-icassp2022.onnx")
parser.add_argument("--data_dir", "-d", required=True, help="Directory containing mic and farend files.")
parser.add_argument("--output_dir", "-o", required=True, help="Output directory to save enhanced output files.")
parser.add_argument("--output_sr", "-osr", default=48_000, help="Sample rate for output files.")
args = parser.parse_args()
if not os.path.exists(args.output_dir):
print(f"Creating output directory: {args.output_dir}")
os.makedirs(args.output_dir)
model_sampling_rate = 16_000
model = DECModel(
model_path=args.model_path,
window_length=0.02,
hop_fraction=0.5,
dft_size=320,
hidden_size=322,
sampling_rate=model_sampling_rate)
mic_paths = glob.glob(os.path.join(args.data_dir, "*_mic.wav"))
for mic_path in tqdm(mic_paths):
basename = os.path.basename(mic_path)
farend_path = mic_path.replace("_mic.wav", "_lpb.wav")
if not os.path.exists(farend_path):
print("Farend file not found, skipping:", farend_path)
continue
out_path = os.path.join(args.output_dir, basename)
if os.path.exists(out_path):
print("Enhanced file exists, overwriting:", out_path)
x_enhanced = model.enhance(mic_path, farend_path)
x_enhanced = librosa.resample(x_enhanced, orig_sr=model_sampling_rate, target_sr=args.output_sr)
sf.write(out_path, x_enhanced, args.output_sr)
二、Pytorch原始模型结构实现
官方给出的是ONNX(Open Neural Network Exchange)格式的模型,主要用于在不同的深度学习框架之间进行模型的转换和部署。它允许您将模型从一个框架(如 PyTorch、TensorFlow 等)导出为 ONNX 格式,然后在支持 ONNX 的其他框架或推理引擎中进行推理。
如果希望对模型进行重新训练,需要使用原始的深度学习框架(如 PyTorch 或 TensorFlow)中的模型定义和训练代码。
- 定义模型 使用 PyTorch 定义模型的结构。
- 准备数据 加载和预处理训练数据。
- 训练模型使用训练数据对模型进行训练,并进行验证。
- 导出为 ONNX 训练完成后,可以将模型导出为 ONNX 格式,以便在其他平台上进行推理。
下面代码提供了一种模型实现:
import torch
import torch.nn as nn
class AECModel(nn.Module):
def __init__(self, input_size=322, hidden_size=322, output_size=161):
super(AECModel, self).__init__()
self.gru1 = nn.GRU(input_size=input_size, hidden_size=hidden_size, batch_first=True)
self.gru2 = nn.GRU(input_size=hidden_size, hidden_size=hidden_size, batch_first=True)
self.fc = nn.Linear(hidden_size, output_size)
def forward(self, x):
out, h01 = self.gru1(x)
out = out + x
out, h02 = self.gru2(out)
out = self.fc(out)
out = torch.sigmoid(out)
out = torch.clamp(out, min=0.0, max=1.0)
return out
# 测试模型
if __name__ == "__main__":
# 模型参数
input_size = 322
hidden_size = 322
output_size = 161
batch_size = 1
num_frames = 10
model = AECModel(input_size, hidden_size, output_size)
input_tensor = torch.randn(batch_size, num_frames, input_size).float()
print(f'Input shape:{input_tensor.shape}')
output = model(input_tensor)
print(f"Final Output shape: {output.shape}")
更多推荐




所有评论(0)