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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 15:17:57