PyTorch优化nn.Embedding时BCELoss.backward报错及训练不收敛求助
问题根因
- 计算图断裂是核心问题:你使用列表推导生成每个边的预测值后,通过
torch.FloatTensor()构造新张量的操作,会直接提取列表中每个张量的纯数值,完全丢弃之前的计算历史,导致生成的res张量和nn.Embedding的参数之间没有计算图关联,反向传播时梯度无法回传到Embedding层,因此会触发无grad_fn的报错。 - 手动给
res添加requires_grad=True无效:这个操作仅能让res自身支持梯度计算,但无法恢复已经断裂的上游计算链,梯度只会停留在res层,不会传递到Embedding参数,因此参数不会更新,loss无下降趋势。
修正方案
优先使用PyTorch向量化操作替代Python循环+列表推导,既保留完整计算图,运行效率也更高:
import torch import torch.nn as nn from torch.optim import SGD emb = nn.Embedding(num_node, embedding_dim) optimizer = SGD(emb.parameters(), lr=0.1, momentum=0.9) loss_fn = nn.BCELoss() sigmoid = nn.Sigmoid() for i in range(500): optimizer.zero_grad() # 向量化取embedding,避免Python循环 u_emb = emb(train_edge[0]) v_emb = emb(train_edge[1]) # 按元素相乘后按行求和,等价于每对节点的点乘 res = sigmoid((u_emb * v_emb).sum(dim=1)) loss = loss_fn(res, train_label) loss.backward() optimizer.step() print(f'loss:{loss.item()}')
补充说明
如果确实需要处理非向量化的特殊逻辑,不要用torch.FloatTensor()封装列表,改用torch.stack()拼接列表中的张量,stack操作会保留每个张量的计算历史,也能解决计算图断裂问题:
res = torch.stack([sigmoid(torch.dot(emb(a), emb(b))) for (a, b) in zip(train_edge[0], train_edge[1])])
另外建议使用loss.item()代替直接打印loss,可以避免每次打印都生成计算图相关的冗余缓存,占用显存。
内容的提问来源于stack exchange,提问作者Vid Stropnik
相关产品推荐
相关产品推荐

