如何在PyTorch中用张量索引子集化张量且避免二次反向传播错误?
PyTorch嵌入张量任意子集化训练问题解决
问题背景
我正在编写基于嵌入计算的PyTorch程序,训练时需要使用嵌入张量的不同子集。用张量索引做子集化会触发二次反向传播错误,但切片操作却能正常运行。切片版本已验证可正常训练,但新功能无法仅通过切片实现。
运行环境:
- Python 3.10.9
- PyTorch 2.0.0(搭配CUDA 11.8)
- Windows系统
复现代码:
import torch device = 'cpu' # 使用device='cuda:0'时错误相同 embeddings = torch.tensor(torch.randn([1024, 128], device=device, dtype=torch.float32), requires_grad=True) target = torch.rand([1024, 128], device=device) optimizer = torch.optim.Adadelta([embeddings], lr=1.0) # 尝试nn.Embedding同样报错 embeddings_2 = torch.nn.Embedding(1024, 128, device=device) # OPTION 1: 正常运行 # cur_embeddings = embeddings[:512, :] # OPTION 2: 需求功能,但报错 cur_embeddings = embeddings[torch.arange(512, device=device), :] # OPTION 3: nn.Embedding方式也报错 # cur_embeddings = embeddings_2(torch.arange(512, device=device)) cur_target = target[:512, :] for idx in range(2): optimizer.zero_grad() loss = torch.nn.MSELoss()(cur_embeddings, cur_target) loss.backward() optimizer.step()
三个选项中,仅选项1正常运行,选项2、3均报错,需实现选项2的任意行子集化功能并保证训练正常。
问题原因
切片属于连续区域选择,PyTorch可直接对原张量的连续块更新梯度;而张量索引是高级索引,会生成原张量的非连续视图或副本,导致Adadelta等优化器维护的动量状态无法正确关联原张量的梯度更新,最终触发二次反向传播错误。
解决方法
方法1:用torch.index_select替代直接索引
torch.index_select会明确返回原张量的视图,确保反向传播时梯度能正确传递到原嵌入张量:
# 替换选项2的代码 indices = torch.arange(512, device=device) cur_embeddings = torch.index_select(embeddings, dim=0, index=indices)
方法2:对nn.Embedding权重使用index_select
若使用nn.Embedding,直接调用层会生成新张量,改为对其权重参数执行index_select即可:
# 替换选项3的代码 indices = torch.arange(512, device=device) cur_embeddings = torch.index_select(embeddings_2.weight, dim=0, index=indices)
注意事项
每次迭代前必须调用optimizer.zero_grad()清空梯度,避免梯度累积导致的冲突。
内容的提问来源于stack exchange,提问作者Ben W
相关产品推荐
相关产品推荐

