如何将多个函数合并为接收shape=(8,)参数并返回shape=(8,)结果的函数
问题描述
我有以下两个函数:
import numpy as np def eq_2(x): A, P, E, EA = x return np.array([E*A, EA, EA, E*P]) def eq_3(x): A, P, E, EA = x return np.array([E**2, E, E, E])
随后我将它们存入列表并命名为v:
v = [eq_2, eq_3] # 输出示例:[<function eq_2 at 0x7f2>, <function eq_3 at 0x7f3>]
我的需求是:如何把v当作一个单一函数来使用,使其接收shape=(8,)的参数x,并返回shape=(8,)的结果?另外,我希望这个方案支持合并任意数量的函数(也就是可以扩展v中的元素数量)。
解决方案
你可以编写一个包装函数,把列表里的多个函数整合起来,按规则拆分输入参数、调用每个子函数,最后拼接结果。具体实现如下:
1. 通用包装函数实现
假设每个子函数都接收shape=(4,)的输入并返回同长度输出,包装函数的逻辑是:
- 把输入的长数组
x按子函数数量拆分成多个子数组 - 依次调用每个子函数,收集返回值
- 拼接所有返回值成一个长数组
代码示例:
def combine_functions(func_list): def combined(x): sub_input_len = 4 # 对应每个子函数接收的输入长度 # 拆分输入为对应数量的子数组 split_x = np.array_split(x, len(func_list)) # 调用每个子函数并收集结果 results = [func(sub_x) for func, sub_x in zip(func_list, split_x)] # 拼接结果 return np.concatenate(results) return combined
2. 使用方式
用包装函数把v转换成单一函数,传入shape=(8,)的参数测试:
# 创建合并后的函数 combined_func = combine_functions(v) # 测试输入:前4个元素给eq_2,后4个给eq_3 test_x = np.array([1, 2, 3, 4, 5, 6, 7, 8]) output = combined_func(test_x) print(output.shape) # 输出: (8,)
3. 扩展支持任意数量函数
如果要添加更多子函数,只需把函数加入列表,同时保证输入x的总长度等于子函数数量 × 单个子函数输入长度即可:
# 新增一个示例函数 def eq_4(x): A, P, E, EA = x return np.array([A+P, E-EA, A*E, P/EA]) # 扩展函数列表 v_extended = [eq_2, eq_3, eq_4] combined_func_extended = combine_functions(v_extended) # 测试输入长度为12(3×4) test_x_extended = np.array([1,2,3,4,5,6,7,8,9,10,11,12]) output_extended = combined_func_extended(test_x_extended) print(output_extended.shape) # 输出: (12,)
注意事项
- 如果子函数的输入长度不一致,可以修改包装函数,用元组列表存储函数和对应输入长度(比如
[(func1, 4), (func2, 5)]),然后按指定长度拆分输入x - 必须保证输入
x的总长度等于所有子函数输入长度的总和,避免拆分出错
内容的提问来源于stack exchange,提问作者David Streuli
相关产品推荐
相关产品推荐

