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

训练对称线性自编码器时出现无规律sqrt(MSE)损失尖峰如何解决

自编码器训练损失无规律尖峰问题

问题介绍

我正在训练一个自编码器,用于学习32个时间步下的位置、速度等32个特征,对应输入为32×32的类图像数据。我搭建了简单的对称线性自编码器,编码器和解码器每层都使用Tanh激活函数。
训练时我仅在输入侧添加了自定义的dropout实现,后续计划替换为nn.Dropout。

问题现象

训练过程中(Batch_Size = 6000),损失函数"sqrt(MSE)"会无规律出现大幅尖峰。
已完成的测试(单测试最多运行1000轮epoch):

  • 调用clip_grad_norm_(model.parameters(), max_norm = 0.5)进行梯度裁剪
  • 尝试替换激活函数为ReLu和ELU
  • 尝试将Batch大小调整为原计划全量的1/2(原计划用全量但GPU显存不足)
  • 关闭输入噪声和dropout(噪声/dropout对问题有缓解但未彻底解决)
  • 取消MSE损失的平方根计算

相关代码

训练核心代码

def rand_bin_array(p_zeros, shape):
    size = 1
    for e in shape:
        size *= e
    arr = np.ones(size)
    arr[:int(size * p_zeros)] = 0
    np.random.shuffle(arr)
    arr = arr.reshape(shape)
    return arr

class Autoencoder_Liniar(nn.Module):
  def __init__(self):
    super().__init__()
    self.encoder = nn.Sequential(
        nn.Linear(1024, 921),
        nn.Tanh(),
        nn.Linear(921, 736),
        nn.Tanh(),
        nn.Linear(736, 515),
        nn.Tanh(),
        nn.Linear(515, 309),
        nn.Tanh(),
        nn.Linear(309, 128),
        nn.Tanh(),
        nn.Linear(128, 64),
        nn.Tanh(),
    )

    self.decoder = nn.Sequential(
        nn.Linear(64, 128),
        nn.Tanh(),
        nn.Linear(128, 309),
        nn.Tanh(),
        nn.Linear(309, 515),
        nn.Tanh(),
        nn.Linear(515, 736),
        nn.Tanh(),
        nn.Linear(736, 921),
        nn.Tanh(),
        nn.Linear(921, 1024),
        nn.Tanh()
    )
  
  def forward(self, x):
    enc = self.encoder(x)
    dec = self.decoder(enc)
    return dec

torch.manual_seed(0)
model = Autoencoder_Liniar().cuda()

criterion = nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

random.seed(0)
epochs = 10000
batch_size = 6000
test_b_size = 5000

train_losses = []
test_losses = []
for i in range(epochs):
  avg_loss = 0
  random.shuffle(train_data)
  for b in range(train_nr // batch_size):
    start = b * batch_size
    data = torch.FloatTensor(train_data[start : start + batch_size]).cuda()

    noise_power = max(0.8 - i/epochs, 0.1)
    noise = torch.FloatTensor(rand_bin_array(noise_power, data.shape)).cuda()
    y_pred = model(data * noise)
    loss = torch.sqrt(criterion(y_pred, data))

    optimizer.zero_grad()
    loss.backward()
    #torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.5)
    optimizer.step()

    avg_loss += loss.item()
    if b % 20 == 0:
      print(f'EPOCH: {i} BATCH: {b} LOSS: {loss.item()}')
  
  train_losses.append(avg_loss / (train_nr // batch_size))

  with torch.no_grad():
    avg_loss = 0
    for b in range(test_nr // test_b_size):
        start = b * test_b_size
        data = np.array(test_data[start : start + test_b_size])
        data = torch.FloatTensor(data).cuda()

        y_pred = model(data)
        loss = torch.sqrt(criterion(y_pred, data))
        avg_loss += loss.item()

  test_losses.append(avg_loss / (test_nr // test_b_size))

梯度统计代码

total_norm = 0
for p in model.parameters():
    param_norm = p.grad.detach().data.norm(2)
    total_norm += param_norm.item() ** 2
total_norm = total_norm ** 0.5
avg_grad += total_norm

optimizer.step()

原因分析

  1. 解码器最后一层用了Tanh激活,输出范围固定为[-1,1],如果输入数据没有归一化到对应区间,模型天生无法拟合超出范围的输入值,遇到分布偏移的batch时误差会突然大幅上升,形成尖峰。
  2. 自定义的dropout实现不符合标准逻辑:你是全局固定比例置0后打乱,不是每个元素独立按概率置0,会导致噪声分布不稳定,极端情况下会触发梯度爆炸。
  3. 训练时直接对MSE损失求平方根,当损失值较小时,平方根运算的梯度会急剧变大,引发梯度爆炸,导致参数更新失控。
  4. 优化器没有设置学习率衰减,固定0.001的学习率在训练后期参数接近收敛时,会导致参数更新幅度过大,损失跳升。
  5. 代码中梯度裁剪功能处于注释状态,没有实际生效,无法抑制偶发的梯度爆炸。

修复方案

  1. 要么将输入数据归一化到[-1,1]区间匹配Tanh输出,要么移除解码器最后一层的Tanh激活,改用线性输出,让模型可以拟合任意范围的输入。
  2. 替换自定义dropout实现为PyTorch官方的nn.Dropout,训练时自动启用,验证时自动关闭,避免自定义实现的分布错误。
  3. 训练阶段直接使用MSE损失进行反向传播,仅在需要统计指标的时候单独计算sqrt(MSE),避免平方根运算带来的梯度异常。
  4. 给优化器添加学习率衰减策略,比如使用ReduceLROnPlateau在验证损失停止下降时自动降低学习率,或者设置固定步长的学习率衰减。
  5. 放开梯度裁剪的注释,实际启用梯度裁剪功能,max_norm设置为0.3~1之间即可,有效抑制偶发的梯度爆炸。
  6. 可在每一层Linear层后添加LayerNorm,稳定训练过程中的特征分布,进一步降低尖峰出现的概率。

内容的提问来源于stack exchange,提问作者Dan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 07:42:02