求可微分带通滤波器实现代码,支持中心频率与带宽微分
可微分带通滤波器实现方案
针对需要对中心频率、带宽支持微分的带通滤波器,以下是几种实用的实现思路和代码示例(基于PyTorch,自动微分框架天然支持梯度计算):
一、FIR可微分带通滤波器(窗函数法)
FIR滤波器结构简单,线性相位特性友好,通过将中心频率(f0)和带宽(bw)作为可训练参数,利用自动微分框架直接计算梯度。
核心原理
基于汉宁窗的带通滤波器系数公式:
h[n] = (2bw/Fs) * sinc(2bw*(n-N/2)/Fs) * cos(2πf0(n-N/2)/Fs)
其中N为滤波器阶数,Fs为采样率,n取0到N-1。
PyTorch代码实现
import torch import numpy as np class DifferentiableFIRBPF(torch.nn.Module): def __init__(self, fs, filter_order, init_f0, init_bw): super().__init__() self.fs = fs self.N = filter_order # 注册可训练参数:中心频率、带宽 self.f0 = torch.nn.Parameter(torch.tensor(init_f0, dtype=torch.float32)) self.bw = torch.nn.Parameter(torch.tensor(init_bw, dtype=torch.float32)) # 生成采样点索引张量 self.n = torch.arange(self.N, dtype=torch.float32) - self.N/2 def forward(self, x): # 计算滤波器系数 sinc_term = torch.sinc(2 * self.bw * self.n / self.fs) cos_term = torch.cos(2 * np.pi * self.f0 * self.n / self.fs) hann_window = torch.hann_window(self.N, dtype=torch.float32) h = (2 * self.bw / self.fs) * sinc_term * cos_term * hann_window h = h / torch.sum(h) # 归一化增益 # 用1D卷积实现滤波 x = torch.nn.functional.conv1d(x.unsqueeze(1), h.unsqueeze(0).unsqueeze(0), padding=self.N//2) return x.squeeze(1)
使用时直接将该模块作为可训练组件加入你的滤波器组,优化器会自动计算f0和bw的梯度。
二、IIR可微分带通滤波器(状态空间形式)
二阶Butterworth带通滤波器的状态空间实现,避免递归循环导致的梯度不稳定,适合需要更窄带宽的场景。
PyTorch代码实现
import torch import numpy as np class DifferentiableIIRBPF(torch.nn.Module): def __init__(self, fs, init_f0, init_bw): super().__init__() self.fs = fs self.f0 = torch.nn.Parameter(torch.tensor(init_f0, dtype=torch.float32)) self.bw = torch.nn.Parameter(torch.tensor(init_bw, dtype=torch.float32)) def get_state_matrix(self): # 计算Butterworth带通的状态空间矩阵 w0 = 2 * np.pi * self.f0 / self.fs alpha = np.pi * self.bw / self.fs cos_w0 = torch.cos(w0) sin_w0 = torch.sin(w0) A = torch.tensor([ [cos_w0, sin_w0], [-sin_w0, cos_w0] ], dtype=torch.float32) * torch.exp(-alpha) B = torch.tensor([[1 - torch.exp(-alpha)], [0]], dtype=torch.float32) C = torch.tensor([[1 - torch.exp(-alpha), 0]], dtype=torch.float32) return A, B, C def forward(self, x): A, B, C = self.get_state_matrix() batch_size, seq_len = x.shape # 初始化状态 state = torch.zeros(batch_size, 2, dtype=torch.float32, device=x.device) output = [] for t in range(seq_len): state = torch.matmul(state, A.T) + torch.matmul(x[:, t].unsqueeze(1), B.T) out = torch.matmul(state, C.T).squeeze(1) output.append(out) return torch.stack(output, dim=1)
三、频域可微分带通滤波器(STFT掩码法)
通过构造可微分的频域掩码,结合STFT/逆STFT实现滤波,适合需要并行处理长序列的场景。
PyTorch代码实现
import torch import torchaudio.transforms as T class DifferentiableFreqBPF(torch.nn.Module): def __init__(self, fs, n_fft=512, hop_length=128, init_f0=None, init_bw=None): super().__init__() self.fs = fs self.n_fft = n_fft self.hop_length = hop_length self.stft = T.Spectrogram(n_fft=n_fft, hop_length=hop_length, power=None) self.istft = T.InverseSpectrogram(n_fft=n_fft, hop_length=hop_length) # 注册可训练参数 self.f0 = torch.nn.Parameter(torch.tensor(init_f0, dtype=torch.float32)) self.bw = torch.nn.Parameter(torch.tensor(init_bw, dtype=torch.float32)) # 生成频率轴 self.freqs = torch.linspace(0, fs/2, n_fft//2 + 1) def forward(self, x): # 计算STFT spec = self.stft(x) # 构造高斯掩码(可微分) mask = torch.exp(-((self.freqs - self.f0)**2) / (2 * self.bw**2)) # 应用掩码 masked_spec = spec * mask.unsqueeze(0).unsqueeze(2) # 逆STFT output = self.istft(masked_spec, length=x.shape[1]) return output
内容的提问来源于stack exchange,提问作者Mason Wang
相关产品推荐
相关产品推荐

