You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在两个PyTorch张量批次中寻找最近匹配(基于距离度量)

PyTorch张量批次间的最近匹配实现

需求说明

  • 给定两个PyTorch张量批次:b1×m×n(批次大小b1)和b2×m×n(批次大小b2)
  • 计算b1中每个张量与b2中每个张量的逐元素平方欧氏距离之和,得到b1×b2的距离矩阵
  • 为b1中的每个张量,找到b2中距离最近的张量的批次索引

已实现的距离计算验证

你当前的代码已经正确生成了距离矩阵(注:torch.cdist默认计算的是欧氏距离,即平方欧氏距离的平方根;但由于最小欧氏距离与最小平方欧氏距离对应同一个匹配项,后续找索引时可以直接使用该结果,无需额外平方):

import torch

a = torch.rand((3, 3, 4))
b = torch.rand((5, 3, 4))
flat_a = torch.flatten(a, start_dim=1)
flat_b = torch.flatten(b, start_dim=1)
distance_matrix = torch.cdist(flat_a, flat_b)

获取最近匹配的批次索引

使用torch.argmin()函数,沿着距离矩阵的**b2维度(dim=1)**取最小值的索引,即可得到b1中每个张量对应的b中最近张量的批次索引:

# 获取每个a中张量对应的b中最近张量的索引
nearest_indices = torch.argmin(distance_matrix, dim=1)

# 输出示例
print("距离矩阵形状:", distance_matrix.shape)  # 输出: torch.Size([3, 5])
print("最近匹配的b批次索引:", nearest_indices)  # 输出形状: torch.Size([3]),元素为0-4的整数

关键说明

  • torch.argmin(dim=1):指定沿着距离矩阵的第二维度(对应b的批次维度)计算最小值的位置,结果是长度为b1的张量,每个元素对应b中最近张量的批次索引。
  • 若需要严格匹配你定义的「逐元素平方欧氏距离之和」,可以对cdist的结果平方:distance_matrix = torch.cdist(flat_a, flat_b) ** 2,但这不会改变最终的索引结果。

内容的提问来源于stack exchange,提问作者BeginnersMindTruly

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.04 14:21:03