PyTorch中如何通过索引优雅获取[N,C]形状的子张量?
解决PyTorch中提取特定位置子张量的优雅方法
针对你遇到的场景,这里有两种无需转换Python列表的优雅实现方式:
方法一:利用张量转置与元组解包
你已经通过torch.nonzero(cond)得到了[N,3]的索引张量,只需将其转置后用tuple()包裹,直接作为索引传入原张量即可:
indices = torch.nonzero(cond) result = x[tuple(indices.T)]
转置后的indices.T是[3,N]的张量,拆分为三个1D张量分别对应B、W、H维度的索引,作为元组传入后,PyTorch会同时对三个维度进行索引,最终得到形状为[N,C]的目标张量,完全基于张量操作,无需转换为Python列表。
方法二:直接使用掩码索引(更简洁高效)
既然已经有条件掩码cond,可以跳过计算indices的步骤,直接利用掩码操作提取目标子张量:
# 若cond是0/1整数张量,先转为布尔型 cond_bool = cond.bool() # 扩展掩码维度以匹配原张量的通道维度,再索引并重塑形状 result = x[cond_bool.unsqueeze(-1)].view(-1, x.shape[-1]) # 或者用torch.masked_select实现 result = torch.masked_select(x, cond.unsqueeze(-1)).view(-1, x.shape[-1])
这种方式无需额外计算索引,代码更简洁,且掩码操作在PyTorch中经过高度优化,执行效率更高。
补充说明:为什么x[indices]不符合预期
当你直接传入2D张量indices作为索引时,PyTorch会将其视为对原张量**第一个维度(B维度)**的批量索引,相当于从B维度中选取N个[W,H,C]的子张量,最终得到形状为[N,W,H,C]的结果,这和你需要的[N,C]不符。而将索引转置后拆分为多维度索引元组,才能实现对(B,W,H)三个维度的同时定位。
内容的提问来源于stack exchange,提问作者Kolya Ivankov
相关产品推荐
相关产品推荐

