import numpy as npimport matplotlib.pyplot as pltplt.rcParams['font.sans-serif'] = ['SimHei','Microsoft YaHei','WenQuanYi Zen Hei']plt.rcParams['axes.unicode_minus'] = Falseplt.rcParams['font.family'] ='sans-serif'
class LMSFilter: """标准LMS自适应滤波器实现。适用于实时或离线处理,用于系统辨识、噪声消除等场景。""" def __init__(self, filter_order, step_size=0.01): """初始化滤波器。 参数: filter_order (int): 滤波器阶数(抽头数)。阶数越高,模型能力越强,但计算量越大,收敛可能越慢。 step_size (float): 步长参数 μ。必须在稳定范围内 (0 < μ < 2/输入功率)。通常从0.01这样的小值开始尝试。 """ self.order = filter_order self.mu = step_size self.weights = np.zeros(filter_order) self.input_buffer = np.zeros(filter_order)
def adapt(self, desired, reference): """处理一个采样点,并更新滤波器。 参数: desired (float): 期望信号 d(n),即主输入(如带噪语音)。 reference (float): 参考输入 x(n),与干扰噪声相关(如参考麦克风采集的噪声)。 返回: output (float): 滤波器输出 y(n)。 error (float): 误差信号 e(n),通常作为系统输出(降噪后的信号)。 """ self.input_buffer[1:] = self.input_buffer[:-1] self.input_buffer[0] = reference output = np.dot(self.weights, self.input_buffer) error = desired - output self.weights += self.mu * error * self.input_buffer return output, error
def filter(self, desired_signal, reference_signal): """批量处理整个信号序列。 参数: desired_signal (np.array): 期望信号序列。 reference_signal (np.array): 参考信号序列。必须与desired_signal长度相同。 返回: output_signal (np.array): 滤波器输出序列(预测的噪声)。 error_signal (np.array): 误差信号序列(估计的干净信号)。 weights_history (np.array): 权重向量的历史记录,用于分析收敛过程。 """ n_samples = len(desired_signal) output_signal = np.zeros(n_samples) error_signal = np.zeros(n_samples) weights_history = np.zeros((n_samples, self.order)) for i in range(n_samples): output_signal[i], error_signal[i] = self.adapt(desired_signal[i], reference_signal[i]) weights_history[i, :] = self.weights return output_signal, error_signal, weights_history
def test_lms_basic(): """基础测试:滤除一个正弦波干扰。""" np.random.seed(42) Fs = 1000 t = np.arange(0, 1.0, 1 / Fs) clean = 0.5 * np.sin(2 * np.pi * 10 * t) noise = 0.8 * np.sin(2 * np.pi * 50 * t) desired = clean + noise reference = 1.2 * noise + 0.05 * np.random.randn(len(t)) lms = LMSFilter(filter_order=32, step_size=0.02) _, error_signal, weights_history = lms.filter(desired, reference) fig, axes = plt.subplots(3, 1, figsize=(10, 8)) axes[0].plot(t, desired, 'b', alpha=0.7, label='带噪信号 (Desired)') axes[0].plot(t, clean, 'r--', linewidth=2, label='原始干净信号 (目标)') axes[0].set_ylabel('幅度') axes[0].set_title('输入信号对比') axes[0].legend() axes[0].grid(True) axes[1].plot(t, error_signal, 'g', label='LMS输出 (Error)') axes[1].plot(t, clean, 'r--', linewidth=2, label='原始干净信号') axes[1].set_ylabel('幅度') axes[1].set_title('降噪效果对比') axes[1].legend() axes[1].grid(True) axes[2].plot(weights_history[:, 5], 'm') axes[2].set_xlabel('采样点') axes[2].set_ylabel('权重值') axes[2].set_title('滤波器权重收敛过程 (示例: w[5])') axes[2].grid(True) plt.tight_layout() plt.show() mse_before = np.mean((desired - clean) ** 2) mse_after = np.mean((error_signal - clean) ** 2) print(f"降噪前MSE:{mse_before:.6f}") print(f"降噪后MSE:{mse_after:.6f}") print(f"MSE改善:{10 * np.log10(mse_before / mse_after):.2f}dB")
if __name__ == "__main__": test_lms_basic()