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

如何编写兼容多数据类型的通用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 01:00:13