如何避免MLP训练初始迭代出现NaN损失与梯度为None问题
解决MLP拟合IIR滤波器系数时梯度为None、损失NaN的问题
问题根源分析
- 模型参数流混淆:
FilterNet的forward方法中用self.sos = self.mlp(x)覆盖了初始化时定义的可训练参数self.sos,冗余参数会干扰优化器的参数更新逻辑,同时易导致梯度流断裂。 - 学习率过高:设置的
lr=1e-1过大,MLP输出的滤波器系数易出现突变,直接引发IIR滤波器不稳定,信号发散产生NaN。 - 滤波器系数无约束:biquad滤波器的分母系数(a组)未做稳定性约束,当系数组合导致极点超出单位圆时,滤波信号会指数增长,触发损失NaN。
- 冗余梯度操作:手动执行
sos.requires_grad_(True)等操作完全多余,MLP输出本身处于计算图中自带梯度属性,手动设置反而可能干扰梯度传递。
修复后的代码
import time import torch import torchaudio import numpy as np from tqdm import tqdm from torchaudio.functional import lfilter from torch.optim import Adam, lr_scheduler # 设备设置 hardware = "cpu" device = torch.device(hardware) class FilterNet(torch.nn.Module): def __init__(self, input_size, hidden_size, output_size, num_biquads=1, fs=44100): super(FilterNet, self).__init__() self.eps = 1e-8 self.fs = fs # 移除冗余的self.sos参数,直接由MLP输出系数 self.mlp = torch.nn.Sequential( torch.nn.Linear(input_size, 100), torch.nn.ReLU(), torch.nn.Linear(100, 50), torch.nn.ReLU(), torch.nn.Linear(50, output_size) ) def get_dirac(self, size, index=0, grad=False): tensor = torch.zeros(size, requires_grad=grad, device=device) tensor[index] = 1 return tensor def compute_filter_magnitude_and_phase_frequency_response(self, dirac, fs, a, b): filtered_dirac = lfilter(dirac, a, b) freqs_response = torch.fft.fft(filtered_dirac) freqs_rad = torch.fft.rfftfreq(filtered_dirac.shape[-1]) freqs_hz = freqs_rad[:filtered_dirac.shape[-1] // 2] * fs / np.pi freqs_response = freqs_response[:len(freqs_hz)] # 添加eps避免log10(0)引发NaN mag_response_db = 20 * torch.log10(torch.abs(freqs_response) + self.eps) phase_response_rad = torch.angle(freqs_response) phase_response_deg = phase_response_rad * 180 / np.pi return freqs_hz, mag_response_db, phase_response_deg def forward(self, x): sos = self.mlp(x) # 约束a0近似为1并归一化,保证滤波器稳定性 sos[:, 3] = torch.clamp(sos[:, 3], min=0.9, max=1.1) sos[:, 3:] = sos[:, 3:] / sos[:, 3:][:, 0:1] return sos # 目标滤波器参数 fs = 2048 num_biquads = 1 num_biquad_coeffs = 6 target_sos = torch.tensor([0.803, -0.132, 0.731, 1.000, -0.426, 0.850], device=device) a = target_sos[3:] b = target_sos[:3] # 生成训练数据(修正超出奈奎斯特频率的问题) import scipy.signal as signal f0 = 20 f1 = 1000 # 修正为采样率一半以内的合理值 t = np.linspace(0, 1, fs, dtype=np.float32) # 缩短时长减少计算量 sine_sweep = signal.chirp(t=t, f0=f0, t1=1, f1=f1, method='logarithmic') white_noise = np.random.normal(scale=5e-2, size=len(t)) noisy_sweep = sine_sweep + white_noise train_input = torch.from_numpy(noisy_sweep.astype(np.float32)).to(device) train_target = lfilter(train_input, a, b) # 优化器与训练参数 n_epochs = 20 batche_size = 1 seq_length = 512 seq_step = 512 model = FilterNet(seq_length, 10*seq_length, 6, num_biquads, fs).to(device) optimizer = Adam(model.parameters(), lr=1e-4, betas=(0.9, 0.999), eps=1e-08) # 降低学习率 scheduler = lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=2) criterion = torch.nn.MSELoss() # 计算目标频率响应 freqs_hz, mag_response_db, phase_response_deg = model.compute_filter_magnitude_and_phase_frequency_response( model.get_dirac(fs, 0, grad=False), fs, a, b ) target_frequency_response = torch.hstack((mag_response_db, phase_response_deg)) # 训练初始化 start_time = time.time() pbar = tqdm(total=n_epochs) loss_history = [] num_sequences = int(train_input.shape[0] / seq_length) # 训练循环 for epoch in range(n_epochs): model.train() total_loss = 0 print(f"\n+ Epoch : {epoch}") for seq_id in range(num_sequences): start_idx = seq_id * seq_step end_idx = start_idx + seq_length input_seq_batch = train_input[start_idx:end_idx].unsqueeze(0) target_seq_batch = train_target[start_idx:end_idx].unsqueeze(0) optimizer.zero_grad() # 预测系数并滤波 sos = model(input_seq_batch) y = lfilter(waveform=input_seq_batch, b_coeffs=sos[:, :3], a_coeffs=sos[:, 3:]) batch_loss = criterion(y, target_seq_batch) # 反向传播与参数更新 batch_loss.backward() optimizer.step() total_loss += batch_loss.item() print(f"|=========> Sequence {seq_id}: Loss = {batch_loss.item():.9f}") # 记录与更新 epoch_loss = total_loss / num_sequences loss_history.append(epoch_loss) print("-"*100) print(f"|=========> epoch_loss = {epoch_loss:.6f}") scheduler.step(epoch_loss) pbar.update(1) print("*"*100) # 结束计时 elapsed_time = time.time() - start_time print(f"\n训练耗时 {elapsed_time:.2f} 秒。") # 输出预测系数 predicted_sos = model(train_input[:seq_length].unsqueeze(0)).detach().squeeze() print("\n目标系数:") print(target_sos.cpu().numpy()) print("\n预测系数:") print(predicted_sos.cpu().numpy())
关键修改说明
- 清理冗余参数:删除初始化时定义的
self.sos,直接由MLP输出系数,避免参数流混乱。 - 保证滤波器稳定性:对输出的a系数做归一化处理,强制
a0=1,同时限制a0的波动范围,防止极点超出单位圆导致信号发散。 - 降低学习率:将学习率从
1e-1调整为1e-4,避免参数更新幅度过大引发NaN。 - 修正数据问题:将扫频上限调整到奈奎斯特频率以内,同时缩短信号时长,减少计算负担。
- 避免对数运算异常:计算幅度响应时添加小常量
eps,防止出现log10(0)的情况。
内容的提问来源于stack exchange,提问作者SuperKogito
相关产品推荐
相关产品推荐

