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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 14:25:23