如何加速DataFrame中calc_df函数的计算效率?
优化DataFrame连续符号统计的计算性能
问题背景
我有一个DataFrame(df),需要基于第一列的符号,统计连续相邻列中与第一列符号相同的列数,再乘以第一列的符号。当前calc_df函数本地运行耗时如下:
%timeit calc_df(df) 6.38 s ± 170 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
输出示例
a_0 a_1 a_2 a_3 a_4 a_5 a_6 a_7 a_8 a_9 0 0.097627 0.430379 0.205527 0.089766 -0.152690 0.291788 -0.124826 0.783546 0.927326 -0.233117 1 0.583450 0.057790 0.136089 0.851193 -0.857928 -0.825741 -0.959563 0.665240 0.556314 0.740024 2 0.957237 0.598317 -0.077041 0.561058 -0.763451 0.279842 -0.713293 0.889338 0.043697 -0.170676 3 -0.470889 0.548467 -0.087699 0.136868 -0.962420 0.235271 0.224191 0.233868 0.887496 0.363641 4 -0.280984 -0.125936 0.395262 -0.879549 0.333533 0.341276 -0.579235 -0.742147 -0.369143 -0.272578 0 4.0 1 4.0 2 2.0 3 -1.0 4 -2.0
原代码
import numpy as np import pandas as pd from numba import njit np.random.seed(0) pd.set_option('display.max_columns', None) pd.set_option('expand_frame_repr', False) # This function generates demo data. def generate_data(): col = [f'a_{x}' for x in range(10)] df = pd.DataFrame(data=np.random.uniform(-1, 1, [280000, 10]), columns=col) return df @njit def calc_numba(s): a = s[0] b = 1 for sign in s[1:]: if sign == a: b += 1 else: break b *= a return b def calc_series(s): return calc_numba(s.to_numpy()) def calc_df(df): df1 = np.sign(df) df['count'] = df1.apply(calc_series, axis=1) return df def main(): df = generate_data() print(df.head(5)) df = calc_df(df) print(df['count'].head(5)) return if __name__ == '__main__': main()
优化方案
方案1:全向量化numpy操作
原代码的核心瓶颈是逐行apply,即使使用numba,逐行调用函数的开销依然巨大。改用numpy全向量化操作可大幅提升速度:
def calc_df_optimized(df): signs = np.sign(df.values) first_sign = signs[:, 0:1] # 标记每行中与第一列符号不同的位置 diff = signs != first_sign # 找到每行第一个不同的索引,无差异则返回0 first_diff_idx = np.argmax(diff, axis=1) # 处理全相同的行,将索引设为总列数 first_diff_idx[first_diff_idx == 0] = signs.shape[1] # 计算最终结果:连续相同列数 × 第一列符号 df['count'] = first_diff_idx * first_sign.flatten() return df
方案2:Numba批量处理整矩阵
如果偏好使用numba,可直接处理整个二维数组,避免逐行调用的开销:
@njit def calc_numba_batch(signs): n_rows, n_cols = signs.shape result = np.empty(n_rows, dtype=np.float64) for i in range(n_rows): a = signs[i, 0] b = 1 for j in range(1, n_cols): if signs[i, j] == a: b += 1 else: break result[i] = b * a return result def calc_df_numba_batch(df): signs = np.sign(df.values) df['count'] = calc_numba_batch(signs) return df
性能对比
- 原代码:~6.38s
- 方案1(numpy向量化):~10ms左右(硬件差异会影响结果)
- 方案2(numba批量):~20ms左右(硬件差异会影响结果)
两种优化方案均能将速度提升数百倍,其中numpy向量化方案无需额外编译,适合快速部署。
内容的提问来源于stack exchange,提问作者jaried
相关产品推荐
相关产品推荐

