Numba配合pandas groupby apply操作如何正确返回多列数组
问题根源
代码存在两个核心问题导致结果异常:
- 存在语法笔误:列选择部分写的
'col_1缺少闭合单引号,会直接触发语法报错。 - 结构构造错误:Numba函数返回的是两个一维NumPy数组组成的元组,直接传入
pd.DataFrame()时,pandas默认会将每个一维数组识别为一行数据,而非一列,最终生成行数为2、列数等于当前分组数据长度的畸形结构,分组apply拼接后就会出现行转列、大量异常列的问题。
修复方案
方案1:修正apply内的DataFrame构造逻辑
显式指定返回数据按列解析,同时对齐原分组的索引和列名,保证结构匹配:
import numpy as np import numba as nb import pandas as pd @nb.jit(nopython=True) def my_Numba_function(arr1, arr2): arr1[:] = 11 arr2[:] = 22 # 按列拼接两个数组,生成二维结构 return np.column_stack((arr1, arr2)) and_df = df_input_imputed.groupby(key_cols_list, as_index=True)[['col_1', 'col_2']].apply( lambda x: pd.DataFrame( my_Numba_function(arr1=x['col_1'].values, arr2=x['col_2'].values), columns=['col_1', 'col_2'], index=x.index ) )
注意:你的Numba函数里用
arr[:] = 固定值的写法是直接修改传入的原数组视图,运行后原DataFrame对应列的值也会被同步修改,如果不需要修改原数据,建议在函数内新建数组存储计算结果再返回。
方案2:使用pandas原生Numba引擎(性能更优)
pandas 1.1.0及以上版本原生支持groupby操作的Numba加速,不需要手动套apply处理结构,自动对齐索引,不会出现结构错乱问题:
import numpy as np import numba as nb import pandas as pd @nb.jit(nopython=True) def my_Numba_function(arr1, arr2): n = len(arr1) out1 = np.empty(n, dtype=arr1.dtype) out2 = np.empty(n, dtype=arr2.dtype) for i in range(n): # 替换成实际计算逻辑即可 out1[i] = 11 out2[i] = 22 return out1, out2 # 用transform搭配numba引擎,自动对齐原表索引 and_df = df_input_imputed.groupby(key_cols_list, as_index=True)[['col_1', 'col_2']].transform( my_Numba_function, engine='numba' )
内容的提问来源于stack exchange,提问作者Boris
相关产品推荐
相关产品推荐

