如何在NumPy中设置全局类型转换规则为'safe'?
NumPy复合赋值类型转换问题的解决方案
首先明确:NumPy没有提供全局修改ufunc默认casting规则的官方方法,复合赋值操作(比如*=)硬编码使用了same_kind的casting规则,无法通过全局配置更改。不过可以通过以下几种方式规避这个问题:
自定义原地乘法函数
自己实现一个函数,封装numpy.multiply的casting='safe'参数,手动完成原地更新:import numpy as np def safe_imul(a, b): # 用safe casting计算结果,再赋值回原数组 result = np.multiply(a, b, casting='safe') a[:] = result # 使用示例 A = np.array([1.5, 2.5], dtype=np.float64) B = np.array([2, 3], dtype=np.int64) safe_imul(A, B)预先统一数组类型
在执行复合赋值前,将B转换为和A相同的类型(float64),这样*=就不会触发类型转换错误:A = np.array([1.5, 2.5], dtype=np.float64) B = np.array([2, 3], dtype=np.int64).astype(A.dtype) A *= B重载ndarray的乘法方法(进阶方案)
如果需要频繁使用,可以子类化np.ndarray并重载__imul__方法,在内部使用casting='safe'的乘法:import numpy as np class SafeCastArray(np.ndarray): def __new__(cls, input_array, dtype=None, order='K'): return np.asarray(input_array, dtype=dtype, order=order).view(cls) def __imul__(self, other): result = np.multiply(self, other, casting='safe') self[:] = result return self # 使用示例 A = SafeCastArray([1.5, 2.5], dtype=np.float64) B = np.array([2, 3], dtype=np.int64) A *= B
内容的提问来源于stack exchange,提问作者Slothscript
相关产品推荐
相关产品推荐

