如何编写兼容多数据类型的通用Python函数,返回与输入同类型结果
如何编写兼容多种数据类型且返回与输入类型一致的Python函数?
要编写能兼容float、Numpy、Pandas等多种数据类型的Python函数,且始终返回与输入参数类型一致的结果,同时计算过程中会包含一个或多个float值。
示例基础函数
def mycalc(x, a=1.0, b=1.0): return a * x + b
(注:此处已简化问题,理想情况下函数可支持多个类似x的输入参数,且已向量化,可处理Numpy数组和Pandas Series)
当前函数的表现
对于Numpy数组和Pandas Series,该函数运行正常,结果的 dtype 由输入参数决定:
import numpy as np x = np.array([1, 2, 3], dtype="float32") print(mycalc(x).dtype) # float32
import pandas as pd x = pd.Series([1.0, 2.0, 3.0], dtype="float32") print(mycalc(x).dtype) # float32
但使用低精度Numpy浮点数时,dtype会被“提升”为float64,推测是公式中的float参数导致:
x = np.float32(1.0) print(mycalc(x).dtype) # float64
理想情况下,函数可支持Python float、Numpy标量、Numpy数组、Pandas Series、Jax数组甚至Sympy符号变量,但不想添加过多语句处理每种情况。
尝试的解决方案及问题
类型转换法(仅支持标量)
该方法适用于Numpy标量,但处理数组或Series时会报错:
def mycalc(x, a=1.0, b=1.0): a = type(x)(a) b = type(x)(b) return a * x + b assert isinstance(mycalc(1.0), float) assert isinstance(mycalc(np.float32(1.0)), np.float32) mycalc(np.array([1, 2, 3], dtype="float32")) # raises TypeError: expected a sequence of integers or a single integer, got '1.0'
装饰器法(仅支持数组和Series)
使用装饰器复制输入参数的方案仅支持Numpy数组和Pandas Series,无法处理Python float或Numpy标量:
import functools def apply_to_pandas(func): @functools.wraps(func) def wrapper_func(x, *args, **kwargs): if isinstance(x, (np.ndarray, list)): out = func(x, *args, **kwargs) else: out = x.copy(deep=False) out[:] = np.apply_along_axis(func, 0, x, *args, **kwargs) return out return wrapper_func @apply_to_pandas def mycalc(x, a=1.0, b=1.0): return a * x + b mycalc(1.0) # TypeError: copy() got an unexpected keyword argument 'deep'
更新:Numpy 2.x的解决方案
在Numpy 2.x版本中,该问题已解决。新版本中,(np.float32(1.0) + 1).dtype == "float32",因此上述原始函数返回结果的dtype与输入x一致。
内容的提问来源于stack exchange,提问作者Bill
相关产品推荐
相关产品推荐

