如何在PyTorch的2D网格坐标张量中高效查找特定坐标?
PyTorch中高效查找2D张量内的特定坐标
直接用PyTorch原生张量操作即可实现高效查找,无需转换为列表或迭代,核心思路是逐行判断坐标是否完全匹配:
单个目标坐标的查找
基于你的示例代码实现:
import torch # 构造示例张量 positions = torch.arange(20).repeat(2).view(-1,2) xy_dst1 = torch.tensor((5,7)) xy_dst2 = torch.tensor((4,5)) # 查找xy_dst1的匹配情况 matches_dst1 = torch.all(positions == xy_dst1, dim=1) indices_dst1 = torch.nonzero(matches_dst1).squeeze() print(indices_dst1) # 输出: tensor([], dtype=torch.int64) # 查找xy_dst2的匹配情况 matches_dst2 = torch.all(positions == xy_dst2, dim=1) indices_dst2 = torch.nonzero(matches_dst2).squeeze() print(indices_dst2) # 输出: tensor([ 2, 12])
步骤说明
positions == xy_dst:生成与positions同形状的布尔张量,每个元素标记对应位置是否等于目标坐标的分量torch.all(..., dim=1):沿着行维度(dim=1)做逻辑与运算,得到1D布尔张量,每个元素表示对应行是否完全匹配目标坐标torch.nonzero().squeeze():提取所有匹配行的索引,squeeze()用于去除多余维度,得到简洁的一维索引张量
批量目标坐标的查找
如果需要同时查找多个目标坐标,可利用PyTorch的广播机制实现:
# 定义多个目标坐标 targets = torch.tensor([(5,7), (4,5)]) # 批量匹配判断:扩展positions维度后与targets做逐元素相等,再逐行验证全匹配 matches_batch = torch.all(positions.unsqueeze(1) == targets, dim=2) # 遍历每个目标,输出对应匹配索引 for idx, target in enumerate(targets): target_indices = torch.nonzero(matches_batch[:, idx]).squeeze() print(f"坐标{target}的匹配索引: {target_indices}")
内容的提问来源于stack exchange,提问作者Tue
相关产品推荐
相关产品推荐

