PyTorch手动计算MNIST图像欧氏距离出现不对称问题求助
问题原因及解决方法
问题原因
问题出在无符号整数溢出:train_data.data存储的是uint8类型(取值范围0-255的无符号整数),当执行s-t时,如果s的像素值小于t,负数结果会触发无符号整数的循环溢出行为——比如5-10会被计算为251(因为uint8无法表示负数,会自动加上256映射到合法区间)。而反过来t-s的结果是5,两者的差值完全不同,后续的范数平方计算自然就不对称了。
替代计算方法
核心思路是先将数据转换为浮点类型,避免整数溢出问题,以下是几种可行的实现:
方法1:手动转换浮点后计算
def wtrain(i,j): s = train_data.data[i].float() # 转为浮点类型,避免溢出 t = train_data.data[j].float() s_flat = torch.flatten(s) t_flat = torch.flatten(t) # 欧氏距离平方等于L2范数的平方,直接计算即可 dist_sq = torch.norm(s_flat - t_flat) ** 2 return dist_sq
方法2:使用PyTorch内置距离函数
import torch.nn.functional as F def wtrain(i,j): s = train_data.data[i].float().flatten() t = train_data.data[j].float().flatten() # pairwise_distance返回L2距离,平方后得到欧氏距离平方 dist_sq = F.pairwise_distance(s, t) ** 2 return dist_sq
方法3:保留归一化操作的版本(如果需要除以1000)
def wtrain(i,j): s = train_data.data[i].float() / 1000 t = train_data.data[j].float() / 1000 s_flat = torch.flatten(s) t_flat = torch.flatten(t) dist_sq = torch.norm(s_flat - t_flat) ** 2 return dist_sq
修改后,wtrain(0,1)和wtrain(1,0)的结果会完全相等,因为浮点类型的减法不会有溢出问题,s-t与t-s是相反数,它们的L2范数平方必然一致。
内容的提问来源于stack exchange,提问作者user135520
相关产品推荐
相关产品推荐

