PyTorch中用3D张量索引选取4D张量对应值的技术问题
PyTorch中用3D坐标张量索引4D张量的解决方案
问题回顾
给定4D张量 possible_values(形状 torch.Size([2, 5, 5, 4])),维度定义为:
- dim 0: batch
- dim 1: x_axis
- dim 2: y_axis
- dim 3: 坐标(x_i,y_j)对应的特征值
同时有3D坐标张量 coordinates(形状 torch.Size([2, 5, 2])),维度定义为:
- dim 0: batch
- dim 1: (x,y)坐标序列
- dim 2: 单个坐标的(x,y)值
需要从每个batch中,选取coordinates指定坐标对应的特征值(即dim3的4个值)。
关键注意点
示例中的坐标是1-based(比如[1,5]),但PyTorch张量采用0-based索引,必须先将坐标转换为0-based,否则会触发索引越界错误。
解决方案代码
import torch # 初始化示例张量 possible_values = torch.randn(2, 5, 5, 4) # [batch, x_axis, y_axis, feature] coordinates = torch.tensor([ [[1,5], [3,3], [2,4], [1,3], [2,3]], [[1,5], [4,3], [2,1], [5,3], [5,3]] ]) # 1. 转换为0-based索引 coords_0based = coordinates - 1 # 2. 拆分x、y坐标分量 x_idx = coords_0based[..., 0] # 形状 [2,5],对应每个batch的x坐标 y_idx = coords_0based[..., 1] # 形状 [2,5],对应每个batch的y坐标 # 3. 生成batch维度的索引,确保坐标与所属batch对应 batch_idx = torch.arange(possible_values.size(0))[:, None].repeat(1, x_idx.size(1)) # 4. 执行高级索引取值 selected_features = possible_values[batch_idx, x_idx, y_idx] # 验证结果形状:预期为 [2,5,4] print(selected_features.size()) # 输出 torch.Size([2, 5, 4])
原理说明
PyTorch的高级索引支持同时使用多个同形状的张量对不同维度进行索引:
batch_idx对应batch维度,确保每个坐标从对应的batch中取值x_idx和y_idx分别对应x_axis和y_axis维度,定位具体坐标位置- 索引后自动保留feature维度(dim3),最终得到每个坐标对应的4维特征值
内容的提问来源于stack exchange,提问作者00sdf0
相关产品推荐
相关产品推荐

