如何在PyTorch中计算指定形状2D与3D张量的对应欧氏距离
批量计算张量对应行的欧氏距离
给定:
- 张量A的形状为
(batch_size, dim) - 张量B的形状为
(batch_size, N, dim)
要计算A中每一行与B对应行内N个向量的欧氏距离,得到形状为(batch_size, N)的结果,可以通过以下两种方式实现(以PyTorch为例):
方法1:直接利用广播计算
通过扩展A的维度,让其与B的维度匹配,再逐元素计算差的平方和,最后开根号:
import torch # 示例张量 batch_size = 2 N = 3 dim = 4 A = torch.randn(batch_size, dim) B = torch.randn(batch_size, N, dim) # 扩展A的维度为(batch_size, 1, dim),实现广播 A_expanded = A.unsqueeze(1) # 计算欧氏距离:先算维度差的平方,求和后开根号 euclidean_dist = torch.sqrt(torch.sum((A_expanded - B)**2, dim=-1)) # 验证结果形状 print(euclidean_dist.shape) # 输出: torch.Size([2, 3])
方法2:利用距离平方展开式(数值更稳定)
欧氏距离的平方可展开为 ||A - B||² = ||A||² + ||B||² - 2*A·B,通过该公式计算能避免减法带来的数值不稳定问题:
import torch # 示例张量 batch_size = 2 N = 3 dim = 4 A = torch.randn(batch_size, dim) B = torch.randn(batch_size, N, dim) # 计算A的L2范数平方,保持维度以便广播 A_norm = torch.sum(A**2, dim=-1, keepdim=True) # shape: (batch_size, 1) # 计算B中每个向量的L2范数平方 B_norm = torch.sum(B**2, dim=-1) # shape: (batch_size, N) # 计算A与B中每个向量的点积 dot_product = torch.bmm(B, A.unsqueeze(-1)).squeeze(-1) # shape: (batch_size, N) # 计算欧氏距离 euclidean_dist = torch.sqrt(A_norm + B_norm - 2 * dot_product) # 验证结果形状 print(euclidean_dist.shape) # 输出: torch.Size([2, 3])
两种方法都能得到目标形状的结果,方法2更适合高维度场景下的数值稳定性需求。
内容的提问来源于stack exchange,提问作者jupyter
相关产品推荐
相关产品推荐

