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

如何避免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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 00:09:51