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

PyTorch向量化计算b,c,h,w张量样本成对距离方法

实现方法

核心是先对齐torch.cdist的输入维度规则,再调用官方优化的算子计算,全程无Python层循环,计算效率最高。

问题原因

torch.cdist的维度规则是:

  • 输入x1形状为(*, P, M)、x2形状为(*, R, M)时,输出形状为(*, P, R)
  • 其中M是单个特征向量的长度,P/R是两组待计算距离的向量数量,*是需要广播对齐的批次维度
    直接传入形状为(b,c,h,w)的4维张量时,算子会误将w识别为特征向量长度M、h识别为向量数量P/R、b,c识别为批次维度,最终输出形状为(b,c,h,h)的位置级距离矩阵,完全不符合样本级距离的计算需求。

正确实现步骤

  1. 将每个样本的非批量维度(通道、高、宽)拉平为一维特征向量,把4维输入转为形状为(样本数, 特征长度)的2维张量,此时每个样本对应一个长度为c*h*w的特征向量
  2. 将拉平后的两个张量直接传入torch.cdist,即可得到形状为(b1, b2)的成对距离矩阵,矩阵中第i行第j列的元素就是第一个输入的第i个样本和第二个输入的第j个样本的标量距离

参考代码

import torch

# 原始批次参数
b, c, h, w = 1000, 128, 28, 28
# 若使用GPU加速,直接在加载数据时把张量放到cuda上
train_batch = torch.randn(b, c, h, w, device="cuda")
test_batch = torch.randn(b, c, h, w, device="cuda")

# 拉平特征维度:从第1维开始到最后一维全部拉平,保留第0维的批量维
train_flat = train_batch.flatten(start_dim=1)  # 输出形状 (1000, 128*28*28) = (1000, 100352)
test_flat = test_batch.flatten(start_dim=1)

# 计算成对L2距离,和需要的||te-tr||逻辑完全一致
# 若要每个测试样本对应一行距离(方便后续取k近邻),把test_flat放在第一个参数位置
dist_matrix = torch.cdist(test_flat, train_flat, p=2)  # 输出形状 (1000, 1000)

对于你给出的b=2的测试用例,上述代码输出的就是形状为(2,2)的矩阵,结构和期望的[[d(t1[0],t2[0]), d(t1[0],t2[1])],[d(t1[1],t2[0]), d(t1[1],t2[1])]]完全一致。

注意事项

  • 不要手动写广播逻辑实现距离计算(比如(test_flat[:,None] - train_flat[None]).norm(dim=-1)),当b=1000、特征维度超过10万时,中间减法结果的张量大小约400GB,会直接显存溢出。torch.cdist内部分块计算的优化可以把显存占用控制在很低的水平,速度也远快于手动实现。
  • 如果训练集样本量太大无法一次性加载进显存,可以分批次加载训练集,每批次拉平后和固定的测试集拉平特征计算距离,动态维护每个测试样本的k个最小距离和对应标签即可,不需要一次性加载全量训练数据。
  • 如果需要使用L1距离、其他p范数距离,只需要修改torch.cdist的p参数即可,默认p=2为欧氏距离。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.02 07:01:08