如何用Decimal对象计算矩阵逆?能否扩展Numpy支持该操作?
基于Decimal类型的NumPy矩阵求逆问题
2015年的相关问题《Matrix inverse with Decimal type NumPy》未得到明确解答,后续我提出的《Is there a way for python to perform a matrix inversion at 500 decimal precision》问题中,hpaulj提供了替代方案建议。
Decimal与NumPy的基础兼容性
Decimal是Python标准库提供的任意精度数值类型,多数NumPy函数可直接对其操作:
多项式求值示例
np.polyval([Decimal(1),Decimal(2)], Decimal(3.1) ) # 输出:Decimal('5.100000000')
Decimal数组的创建
可以直接转换为NumPy的object类型数组,或初始化后赋值:
# 直接转换 np.array([[Decimal(1),Decimal(2)],[Decimal(3),Decimal(4)]]) # 输出: # array([[Decimal('1'), Decimal('2')], # [Decimal('3'), Decimal('4')]], dtype=object) # 初始化后赋值 matrix_m=np.zeros((2,2), dtype=object) for ix in range(0,2): for iy in range(0,2): matrix_m[ix,iy]=Decimal(ix)+Decimal(iy); # 输出: # array([[Decimal('0'), Decimal('1')], # [Decimal('1'), Decimal('2')]], dtype=object)
部分NumPy函数支持Decimal数组
np.exp、np.sqrt等函数能正常处理Decimal元素的数组:
np.exp( np.array([[Decimal(1),Decimal(2)],[Decimal(3),Decimal(4)]]) ) # 输出: # array([[Decimal('2.718281828'), Decimal('7.389056099')], # [Decimal('20.08553692'), Decimal('54.59815003')]], dtype=object) np.sqrt( np.array([[Decimal(1),Decimal(2)],[Decimal(3),Decimal(4)]]) ) # 输出: # array([[Decimal('1'), Decimal('1.414213562')], # [Decimal('1.732050808'), Decimal('2')]], dtype=object)
与Decimal原生计算结果一致
单个元素的NumPy计算结果和Decimal原生函数输出完全匹配:
np.exp(Decimal(1))==Decimal(1).exp() # 输出:True
自定义高精度常量
借助Decimal可实现高精度常量计算,例如自定义π的计算函数:
def pi(): """Compute Pi to the current precision.""" getcontext().prec += 2 # 中间步骤保留额外精度 three = Decimal(3) lasts, t, s, n, na, d, da = 0, three, 3, 1, 0, 0, 24 while s != lasts: lasts = s n, na = n+na, na+8 d, da = d+da, da+32 t = (t * n) / d s += t getcontext().prec -= 2 return +s # 应用设定的精度 # 调用示例:print(pi()) 输出 3.141592653589793238462643383
NumPy线性代数函数的局限性
NumPy的np.linalg.det和np.linalg.inv无法处理Decimal类型矩阵,执行时会触发错误:
报错示例
# 计算行列式报错 np.linalg.det(np.array([[Decimal(1),Decimal(2)],[Decimal(1),Decimal(3)]])) # 报错堆栈: # File <__array_function__ internals>:180, in det(*args, **kwargs) # File ~\anaconda3\lib\site-packages\numpy\linalg\linalg.py:2154, in det(a) # 2152 t, result_t = _commonType(a) # 2153 signature = 'D->D' if isComplexType(t) else 'd->d' # -> 2154 r = _umath_linalg.det(a, signature=signature) # 2155 r = r.astype(result_t, copy=False) # 2156 return r # 矩阵求逆报错 np.linalg.inv(np.array([[Decimal(1),Decimal(2)],[Decimal(1),Decimal(3)]])) # 报错堆栈: # File <__array_function__ internals>:180, in inv(*args, **kwargs) # File ~\anaconda3\lib\site-packages\numpy\linalg\linalg.py:552, in inv(a) # 550 signature = 'D->D' if isComplexType(t) else 'd->d' # 551 extobj = get_linalg_error_extobj(_raise_linalgerror_singular) # -> 552 ainv = _umath_linalg.inv(a, signature=signature, extobj=extobj) # 553 return wrap(ainv.astype(result_t, copy=False))
错误信息
UFuncTypeError: Cannot cast ufunc 'inv' input from dtype('O') to dtype('float64') with casting rule 'same_kind'
hpaulj提供的替代方案
hpaulj建议将Decimal对象转换为mpmath的mpf类型,再通过mpmath完成矩阵求逆:
# 转换为mpmath矩阵 mp.matrix( np.array([[Decimal(1),Decimal(2)],[Decimal(3),Decimal(4)]]) ) # 输出: # matrix( # [['1.0', '2.0'], # ['3.0', '4.0']]) # 查看元素类型 mp.matrix( np.array([[Decimal(1),Decimal(2)],[Decimal(3),Decimal(4)]]) )[0,0] # 输出:mpf('1.0') # 矩阵求逆 mp.matrix( np.array([[Decimal(1),Decimal(2)],[Decimal(3),Decimal(4)]]) ) **(-1) # 输出: # matrix( # [['-2.0', '1.0'], # ['1.5', '-0.5']])
但该方案存在明显缺陷:丢失Decimal库的特性,需要在mpmath、NumPy和Decimal间频繁转换,且mpf对象的计算速度远慢于Decimal对象。
问题
- 是否有简便方法修改或扩展NumPy代码,使
np.linalg.inv()可以处理Decimal数组? - 是否存在直接使用Decimal对象计算矩阵逆的方法?
内容的提问来源于stack exchange,提问作者ShoutOutAndCalculate
相关产品推荐
相关产品推荐

