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)的位置级距离矩阵,完全不符合样本级距离的计算需求。
正确实现步骤
- 将每个样本的非批量维度(通道、高、宽)拉平为一维特征向量,把4维输入转为形状为
(样本数, 特征长度)的2维张量,此时每个样本对应一个长度为c*h*w的特征向量 - 将拉平后的两个张量直接传入
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
相关产品推荐
相关产品推荐

