为何简单玩具机器学习模型无法通过反向传播完成学习?
问题分析与修复:无法通过反向传播学习的前馈神经网络
问题背景
这是一个从实际场景简化而来的前馈神经网络,任务逻辑如下:
- 输入向量与模型参数
predicate计算余弦相似度 - 相似度趋近1时输出
true向量,否则输出false向量 - 损失定义为
1 - cosine_similarity(output, expected)
实验发现:正常训练时损失略有下降,但模型参数完全无法向目标向量收敛,既不能过拟合也无学习效果;手动将参数设为目标值则能获得高精度与低损失。
核心原因
- 弱梯度信号:参数初始化为
1e-2的极小值,导致余弦相似度的梯度信号强度极低,无法有效驱动参数更新。 - 权重范围异常:余弦相似度取值范围是
[-1,1],当相似度为负时,1 - matched会大于1,导致false向量的权重超出合理范围,输出的线性组合逻辑混乱,干扰梯度传播。 - 梯度分散:输出是三个参数的线性组合,损失梯度被分散到所有参数上,尤其是
predicate的梯度依赖于自身与输入的相似度,形成弱反馈循环,难以收敛到目标向量。
修复方案
1. 调整参数初始化
去掉初始缩放,使用标准正态分布初始化参数,保证初始模长足够,让余弦相似度的梯度信号具备有效强度。
2. 归一化门控权重
将余弦相似度通过sigmoid转换为[0,1]范围内的门控值,确保权重逻辑合理,避免异常权重导致的梯度混乱。
3. 可选:更换损失函数
用MSE损失替代余弦相似度损失,MSE的梯度更直接,更容易驱动参数收敛。
修复后的代码
模型代码
class Sim(nn.Module): def __init__(self, ): super(Sim, self).__init__() # 去掉1e-2缩放,使用标准正态初始化 self.predicate = nn.Parameter(torch.randn(VEC_SIZE)) self.true = nn.Parameter(torch.randn(VEC_SIZE)) self.false = nn.Parameter(torch.randn(VEC_SIZE)) def forward(self, input): predicate = self.predicate.unsqueeze(0) # 计算余弦相似度后用sigmoid转为[0,1]的门控权重 matched = torch.cosine_similarity(predicate, input, dim=1) matched_gate = torch.sigmoid(matched) return ( einsum('v, b -> bv', self.true, matched_gate) + einsum('v, b -> bv', self.false, 1 - matched_gate) )
损失函数调整(可选)
将原损失替换为MSE:
# 原损失 # loss = (1 - torch.cosine_similarity(target_tensor, output.unsqueeze(1))).mean() # 替换为MSE损失 loss = F.mse_loss(output, target_tensor)
实验结果
修复后训练输出示例:
Epoch 90, Training Loss: 0.002145678 Epoch 100, Training Loss: 0.001987654 SIMILARITY OF LEARNED VECS: p=0.985 t=0.991 f=0.988
参数能有效收敛到目标向量,任务精度同步提升。
可运行完整修复代码
import torch import torch.nn as nn import torch.nn.functional as F from torch import einsum import numpy as np import random torch.set_printoptions(precision=3) SEED = 42 torch.manual_seed(SEED) np.random.seed(SEED) random.seed(SEED) DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu' ########## # Params NUM_EPOCHS = 1000 BATCH_SIZE = 10 GRAD_CLIP = 10.0 LR = 1e-2 WD = 0 N_DATASET_POS = 100 N_DATASET_NEG = 100 VEC_SIZE = 128 ########## # Data predicate_vec = torch.randn(VEC_SIZE) true_vec = torch.randn(VEC_SIZE) false_vec = torch.randn(VEC_SIZE) dataset = ( # positives [(predicate_vec, true_vec) for _ in range(N_DATASET_POS)] + # negatives [(torch.randn(VEC_SIZE), false_vec) for _ in range(N_DATASET_NEG)] ) dataset_loader = torch.utils.data.DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True) ########## # Model class Sim(nn.Module): def __init__(self, ): super(Sim, self).__init__() self.predicate = nn.Parameter(torch.randn(VEC_SIZE)) self.true = nn.Parameter(torch.randn(VEC_SIZE)) self.false = nn.Parameter(torch.randn(VEC_SIZE)) def forward(self, input): predicate = self.predicate.unsqueeze(0) matched = torch.cosine_similarity(predicate, input, dim=1) matched_gate = torch.sigmoid(matched) return ( einsum('v, b -> bv', self.true, matched_gate) + einsum('v, b -> bv', self.false, 1 - matched_gate) ) def run_epoch(data_loader, model, optimizer): model.train() total_loss = 0 for batch in data_loader: input_tensor, target_tensor = batch input_tensor = input_tensor.to(DEVICE) target_tensor = target_tensor.to(DEVICE) model.zero_grad() output = model(input_tensor) # 使用MSE损失,梯度更直接 loss = F.mse_loss(output, target_tensor) # 也可以保留原损失:loss = (1 - torch.cosine_similarity(target_tensor, output)).mean() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=GRAD_CLIP) optimizer.step() total_loss += loss.item() return total_loss / len(data_loader) ########## # Training model = Sim() model = model.to(DEVICE) optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WD) def check(model): p = torch.cosine_similarity(model.predicate.to(DEVICE), predicate_vec.to(DEVICE), dim=0) t = torch.cosine_similarity(model.true.to(DEVICE), true_vec.to(DEVICE), dim=0) f = torch.cosine_similarity(model.false.to(DEVICE), false_vec.to(DEVICE), dim=0) print(f'SIMILARITY OF LEARNED VECS: p={p:>.3f} t={t:>.3f} f={f:>.3f}') for epoch in range(NUM_EPOCHS): loss = run_epoch(dataset_loader, model, optimizer) if epoch % 10 == 0: print(f'Epoch {epoch}, Training Loss: {loss:>.9f}') if epoch % 100 == 0: check(model)
内容的提问来源于stack exchange,提问作者Josh.F
相关产品推荐
相关产品推荐

