bfloat16类型PyTorch张量:torch.dot与torch.inner结果差异及选型疑问
bfloat16张量内积计算:torch.dot与torch.inner的差异及正确选择
问题根源:bfloat16精度限制下的实现差异
在CPU环境的bfloat16计算中,torch.dot(包括等价的torch.matmul(a,b)、a @ b)与torch.inner的结果差异,源于两者累加过程的实现逻辑不同:
torch.dot在CPU上处理bfloat16一维张量时,采用了不符合预期的低精度累加优化,导致大量小数值累加时出现严重的精度丢失,甚至量级错误;torch.inner(以及等价的(a*b).sum(-1)、torch.mul(a,b).sum(-1)等)是先逐元素相乘,再对结果求和,累加过程的精度控制更符合内积的数学定义,结果更接近真实值。
从测试结果能明显看出:bfloat16下torch.dot输出的256.完全偏离正确量级,而torch.inner输出的2464.与float下的正确结果(约2477)更接近,误差仅源于bfloat16本身的精度限制。
正确选择:优先使用torch.inner或逐元素乘累加方式
当张量类型为bfloat16时,必须避免使用torch.dot、一维直接torch.matmul、@运算符,这些方法在CPU bf16环境下的实现存在精度问题。
可靠的内积计算方式包括:
torch.inner(a, b)(a * b).sum(-1)torch.mul(a, b).sum(-1)torch.matmul(a.unsqueeze(0), b.unsqueeze(-1)).squeeze()
这些方式的计算逻辑一致,都能得到符合预期的内积结果。
补充:float下差异小的原因
float(32位浮点数)拥有23位尾数,精度远高于bfloat16的8位尾数。即使两种方法的累加逻辑有细微差异,也只会产生极小的舍入误差,结果仍会保持一致(torch.isclose返回True)。
内容的提问来源于stack exchange,提问作者postylem
相关产品推荐
相关产品推荐

