You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

求可微分带通滤波器实现代码,支持中心频率与带宽微分

可微分带通滤波器实现方案

针对需要对中心频率、带宽支持微分的带通滤波器,以下是几种实用的实现思路和代码示例(基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.25 09:07:42