如何为PyTorch张量的多个维度选取特定索引以实现部分张量相加?
如何为PyTorch张量的多个维度选取特定索引以实现部分张量相加?
嘿,我完全懂你这个需求——就是要把y精准地加到x里指定batch和channel对应的区域,对吧?毕竟x是四维的[batch, channel, H, W]张量,你已经选好了特定的batch索引和channel索引,y的形状又刚好对应这些选中的子集,接下来就看怎么正确索引到x的对应位置完成相加。
我给你两种靠谱的实现方式,都是PyTorch里常用的高级索引技巧:
方法一:手动扩展索引形状
这种方式比较直观,就是把batch和channel的索引扩展成和y前两个维度匹配的形状,让PyTorch能精准定位到要相加的区域:
import torch x = torch.randn([10, 7, 128, 128]) batch_idx = torch.tensor([1,3], dtype=torch.int64) channel_idx = torch.tensor([2,3,5], dtype=torch.int64) y = torch.randn([2, 3, 128, 128]) # 把batch索引扩展:每个batch对应所有选中的channel,形状变成[2, 3] batch_expanded = batch_idx.unsqueeze(1).repeat(1, len(channel_idx)) # 把channel索引扩展:每个batch都对应同样的channel集合,形状也变成[2, 3] channel_expanded = channel_idx.unsqueeze(0).repeat(len(batch_idx), 1) # 直接定位到x的对应位置,把y加进去 x[batch_expanded, channel_expanded] += y
方法二:用meshgrid生成索引网格
这种方式更简洁,利用torch.meshgrid直接生成batch和channel的索引组合,省去手动扩展的步骤:
import torch x = torch.randn([10, 7, 128, 128]) batch_idx = torch.tensor([1,3], dtype=torch.int64) channel_idx = torch.tensor([2,3,5], dtype=torch.int64) y = torch.randn([2, 3, 128, 128]) # 生成对应索引网格,indexing='ij'确保是行优先的匹配(每个batch对应所有channel) batch_grid, channel_grid = torch.meshgrid(batch_idx, channel_idx, indexing='ij') # 直接索引相加 x[batch_grid, channel_grid] += y
小验证技巧
你可以打印一下索引后的形状,确认和y的形状一致,避免形状不匹配的报错:
print(x[batch_grid, channel_grid].shape) # 应该输出 torch.Size([2, 3, 128, 128])
这里要注意,PyTorch的高级索引会自动处理后面的H和W维度——当你指定了前两个维度的索引后,后面的所有维度会被默认全部选中,所以不用额外写代码去处理128x128的部分,非常省心~
备注:内容来源于stack exchange,提问作者Cloudy
相关产品推荐
相关产品推荐

