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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 08:00:20