基于NumPy实现扁平可变长度One-Hot编码到函数调用的转换
基于NumPy实现布尔数组到函数调用的映射方案
问题背景
- 给定列表
l = [(4, func1), (3, func2), (6, func3)],元组首元素代表对应函数的参数取值数量(例如func1接受0-3的整数参数) - 该列表已被转换为总长度13的扁平布尔数组(非标准One-Hot编码):数组中唯一的
1对应的位置,对应要调用的函数及参数。比如数组第二个位置为1时调用func1(1),最后一个位置为1时调用func3(5) - 需求:无需显式循环,用NumPy实现高效优雅的转换逻辑,从布尔数组得到对应的函数调用结果
实现方案
以下是基于NumPy的无循环实现:
import numpy as np # 示例业务函数(可替换为实际需求函数) def func1(x): return f"func1 called with parameter {x}" def func2(x): return f"func2 called with parameter {x}" def func3(x): return f"func3 called with parameter {x}" l = [(4, func1), (3, func2), (6, func3)] # 预处理:生成函数区间分界点与函数列表 counts = np.array([item[0] for item in l]) funcs = [item[1] for item in l] # 计算累计和,得到每个函数对应的区间结束位置 cumulative_counts = np.cumsum(counts) def one_hot_to_function_call(one_hot_arr): # 快速定位布尔数组中1的位置 target_pos = np.argmax(one_hot_arr) # 找到该位置所属的函数索引 func_index = np.searchsorted(cumulative_counts, target_pos, side='right') # 计算对应函数的参数值 param = target_pos - (cumulative_counts[func_index-1] if func_index > 0 else 0) # 调用函数并返回结果 return funcs[func_index](param) # 测试用例 # 测试func1(1)对应的布尔数组 test_case1 = np.zeros(13, dtype=bool) test_case1[1] = True print(one_hot_to_function_call(test_case1)) # 输出: func1 called with parameter 1 # 测试func3(5)对应的布尔数组 test_case2 = np.zeros(13, dtype=bool) test_case2[-1] = True print(one_hot_to_function_call(test_case2)) # 输出: func3 called with parameter 5
逻辑说明
- 预处理阶段:
- 提取各函数的参数数量到
counts数组,提取函数对象到funcs列表 - 用
np.cumsum计算累计和,得到每个函数对应的区间结束位置(示例中cumulative_counts为[4,7,13],对应func1覆盖0-3位,func2覆盖4-6位,func3覆盖7-12位)
- 提取各函数的参数数量到
- 转换阶段:
- 用
np.argmax快速定位布尔数组中1的位置(因是类One-Hot结构,唯一的1会被精准匹配) - 用
np.searchsorted直接找到位置所属的函数索引,无需循环遍历区间 - 计算参数值:当前位置减去前一个区间的结束位置(第一个函数直接用目标位置作为参数,因为参数从0开始计数)
- 调用对应函数并返回结果
- 用
该方案全程使用NumPy内置函数替代显式循环,兼顾了执行效率与代码简洁性。
内容的提问来源于stack exchange,提问作者Technaton
相关产品推荐
相关产品推荐

