numpy.einsum有时忽略dtype参数的技术问题咨询
我明白你的问题了——你想在einsum计算时全程用int64精度避免溢出,但又不想把原有的int8数组整个转成int64类型,而且发现dtype参数有时候没达到预期效果。让我来拆解一下这个问题并给出解决方案:
问题根源
首先要搞清楚为什么会出现溢出,以及dtype参数的局限:
- numpy的einsum默认会根据输入数组的 dtype 来推断计算过程中使用的精度。当输入是int8时,乘法和求和都会在int8的范围内执行,一旦中间结果超出int8的范围(-128到127),就会发生溢出,导致数据丢失。
- 你提到的
dtype参数,默认情况下主要是控制输出结果的 dtype,而不是强制中间计算过程使用该类型。如果中间计算已经在低精度下溢出了,即使最后转成int64,也无法恢复丢失的数据——这就是为什么你觉得dtype参数有时候不生效。
解决方案:临时提升计算精度(不修改原数组)
我们可以在einsum的计算过程中,临时将输入数组提升为int64类型,但原数组依然保持int8不变。这种方式既满足了全程用int64计算的需求,又不会改变原数组的存储类型。
示例代码
先看你的原始溢出案例:
import numpy as np A = np.array([[123, 45],[67,89]], dtype='int8') # 溢出的结果,因为全程用int8计算 result_overflow = np.einsum(A, [0,1], A, [0,1], [1]) print(result_overflow) # 输出: array([-94, -38], dtype=int8)
修正后的代码(临时提升精度):
# 临时将A转换为int64传入einsum,原数组A仍为int8 result_correct = np.einsum('ij,ij->j', A.astype(np.int64), A.astype(np.int64)) # 或者用你习惯的旧版调用方式: # result_correct = np.einsum(A.astype(np.int64), [0,1], A.astype(np.int64), [0,1], [1], dtype=np.int64) print(result_correct) # 输出: array([19618, 9946], dtype=int64) print(A.dtype) # 原数组依然是int8,验证:int8
为什么这个方法有效?
A.astype(np.int64)会创建一个临时的int64数组副本,用于einsum的计算,原数组A的存储类型和数据完全不变。- 整个计算过程(乘法、求和)都会在int64的精度下执行,彻底避免了溢出问题,最后输出的结果也是正确的int64类型。
额外说明
如果你担心临时转换数组带来的内存开销,其实numpy的astype操作对于小数组来说可以忽略不计;如果是超大数组,也可以考虑分块处理,但一般情况下直接临时转换是最简洁高效的方案。
内容的提问来源于stack exchange,提问作者ea1
相关产品推荐
相关产品推荐

