在这篇博客中,我们将深入探讨 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。
  • h02: FLOAT[1, 1, 322]

    • 这是第二个 GRU 层的初始隐藏状态,形状同样为 [1, 1, 322]。

1.1.2 模型输出

  • output: FLOAT[0, 0, 161]

    • 这是模型的主要输出,形状为 [0, 0, 161],表示输出序列的长度和特征维度。161 是每一帧用于重建语音信号的频点数量。
  • hn1: FLOAT[1, 0, 322]

    • 这是第一个 GRU 层的最终隐藏状态,形状为 [1, 0, 322],表示在处理完输入后,隐藏状态的维度。
  • hn2: FLOAT[1, 0, 322]

    • 这是第二个 GRU 层的最终隐藏状态,形状同样为 [1, 0, 322]。

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
  1. GRU 层 (GRU_0)

    • 输入:
      • input: (1, 1, 322)
      • h01: (1, 1, 322)
    • 输出:
      • hn1: (1, 1, 322)(最终隐藏状态)
      • output: (1, 1, 322)(GRU 的输出)
      • 节点标识符: 34(GRU_0 的输出张量名称)
  2. Squeeze 层 (Squeeze_1)

    • 输入:
      • 34(GRU_0 的输出,形状为 (1, 1, 322))
    • 输出:
      • 36: (1, 322)(去掉了序列长度为 1 的维度)
      • 节点标识符: 36(Squeeze_1 的输出张量名称)
  3. Add 层 (Add_2)

    • 输入:
      • input: (1, 1, 322)
      • 36: (1, 322)
    • 输出:
      • 37: (1, 1, 322)(加法的结果)
      • 节点标识符: 37(Add_2 的输出张量名称)
  4. GRU 层 (GRU_3)

    • 输入:
      • 37: (1, 1, 322)
      • h02: (1, 1, 322)
    • 输出:
      • hn2: (1, 1, 322)(最终隐藏状态)
      • output: (1, 1, 322)(GRU 的输出)
      • 节点标识符: 59(GRU_3 的输出张量名称)
  5. Squeeze 层 (Squeeze_4)

    • 输入:
      • 59(GRU_3 的输出,形状为 (1, 1, 322))
    • 输出:
      • 61: (1, 322)(去掉了序列长度为 1 的维度)
      • 节点标识符: 61(Squeeze_4 的输出张量名称)
  6. MatMul 层 (MatMul_5)

    • 输入:
      • 61: (1, 322)
      • 假设权重矩阵形状为 (322, 161)
    • 输出:
      • 63: (1, 161)(矩阵乘法的结果)
      • 节点标识符: 63(MatMul_5 的输出张量名称)
  7. Add 层 (Add_6)

    • 输入:
      • 63: (1, 161)
      • 假设偏置形状为 (161)
    • 输出:
      • 64: (1, 161)(加法的结果)
      • 节点标识符: 64(Add_6 的输出张量名称)
  8. Sigmoid 层 (Sigmoid_7)

    • 输入:
      • 64: (1, 161)
    • 输出:
      • 65: (1, 161)(经过 Sigmoid 激活函数后的结果)
      • 节点标识符: 65(Sigmoid_7 的输出张量名称)
  9. 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 处理流程

  1. 加载模型和音频文件。
  2. 对每个音频文件进行增强处理。
  3. 将增强后的音频文件保存到指定目录。

完整代码

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)中的模型定义和训练代码。

  1. 定义模型 使用 PyTorch 定义模型的结构。
  2. 准备数据 加载和预处理训练数据。
  3. 训练模型使用训练数据对模型进行训练,并进行验证。
  4. 导出为 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}")

更多推荐