PyTorch中如何用index张量在dim=1维度替换source张量?
解决PyTorch中scatter方法的维度匹配与替换问题
问题分析
你需要利用index张量(形状[20, 3001])在dim=1维度上,将source张量(形状[20, 3, 3001])的对应值替换到目标张量的指定位置。核心问题是scatter方法要求索引张量与源张量的维度匹配,需要对index做维度扩展。
解决方案代码
import torch # 模拟你的输入张量 index = torch.randint(0, 3, (20, 3001)) # 生成0/1/2的索引,与你的index形状一致 source = torch.randn(20, 3, 3001) # 模拟你的source张量 # 初始化目标张量(与source同形状,初始值可按需调整) target = torch.zeros_like(source) # 关键步骤:将index扩展一个维度,使其形状变为[20, 1, 3001],与source的dim=1维度匹配 expanded_index = index.unsqueeze(1) # 从source中取出index指定通道的值,再用scatter_替换到target的对应位置 selected_values = source.gather(dim=1, index=expanded_index) target.scatter_(dim=1, index=expanded_index, src=selected_values)
代码解释
- 维度扩展:
index.unsqueeze(1)将原形状[20, 3001]变为[20, 1, 3001],这样才能与source的[20, 3, 3001]在dim=1维度上匹配,满足scatter的维度要求。 - 提取对应值:
source.gather(dim=1, index=expanded_index)会根据index取出每个(batch, time_step)对应的通道值,得到形状[20, 1, 3001]的张量。 - 执行替换:
target.scatter_会将提取出的值,放到target中dim=1维度对应index指定的位置,其他位置保持初始值(这里是0)。
验证结果
可以通过以下代码验证替换是否正确:
# 随机选一个batch和time_step检查 batch_idx = 5 time_step = 1000 selected_channel = index[batch_idx, time_step] # 检查target对应位置的值是否与source一致 assert torch.allclose(target[batch_idx, selected_channel, time_step], source[batch_idx, selected_channel, time_step]) # 检查其他通道是否保持初始值0 for channel in range(3): if channel != selected_channel: assert target[batch_idx, channel, time_step] == 0
内容的提问来源于stack exchange,提问作者Javad
相关产品推荐
相关产品推荐

