如何加速pandas.where结合groupby的运算?
优化方案
原代码速度慢的核心原因是df1.where(df2 == 2)会生成一个和df1同样大小的巨型DataFrame,包含大量NaN,既占用内存又拖慢后续的groupby计算。以下是几种针对性的优化思路:
方法一:numpy向量化运算(最推荐)
利用numpy的矩阵运算和分组求和,避免生成中间大DataFrame,大幅提升效率:
import pandas as pd import numpy as np # 1. 对齐mask的行到df1的索引 mask_df2 = df2.loc[df1.index] == 2 mask_arr = mask_df2.values # 形状(34, 3467) # 2. 匹配df1列的n1到mask的列索引 n1_col = df1.columns.get_level_values('n1') n1_indices = df2.columns.get_indexer(n1_col) # 每个df1列对应的mask列索引 # 3. 过滤掉df1中不存在于df2的n1列 valid_cols = n1_indices != -1 df1_valid = df1.loc[:, valid_cols] valid_n1_indices = n1_indices[valid_cols] # 4. 构建对应mask矩阵,与df1_valid的元素相乘 mask_matrix = mask_arr[:, valid_n1_indices] df1_masked = df1_valid.values * mask_matrix # 5. 按n2分组求和 n2_valid = df1_valid.columns.get_level_values('n2').to_numpy() unique_n2, group_idx = np.unique(n2_valid, return_inverse=True) # 用bincount实现高效分组求和 result_arr = np.zeros((df1.shape[0], len(unique_n2))) for row_idx in range(df1.shape[0]): result_arr[row_idx] = np.bincount(group_idx, weights=df1_masked[row_idx]) # 转换为最终DataFrame result = pd.DataFrame(result_arr, index=df1.index, columns=unique_n2)
进一步加速:用numba优化行循环
如果行数较多,可借助numba编译循环,进一步提升速度:
from numba import jit @jit(nopython=True) def fast_group_sum(arr, group_idx, num_groups): result = np.zeros((arr.shape[0], num_groups)) for i in range(arr.shape[0]): for j in range(arr.shape[1]): result[i, group_idx[j]] += arr[i, j] return result # 替换之前的行循环 result_arr = fast_group_sum(df1_masked, group_idx, len(unique_n2)) result = pd.DataFrame(result_arr, index=df1.index, columns=unique_n2)
方法二:分n1批量处理
针对df2的列数(3467)远小于df1的列数,按n1分组处理,减少内存占用:
# 1. 对齐mask的行到df1的索引 mask_df2 = df2.loc[df1.index] == 2 # 2. 初始化结果DataFrame unique_n2 = df1.columns.get_level_values('n2').unique() result = pd.DataFrame(0, index=df1.index, columns=unique_n2) # 3. 遍历每个存在于df1的n1,批量计算求和 for n1 in mask_df2.columns: if n1 not in df1.columns.get_level_values('n1'): continue # 获取df1中当前n1对应的所有列(列是n2) df1_n1 = df1.xs(n1, level='n1', axis=1) # 获取当前n1的mask,广播为列维度 n1_mask = mask_df2[n1].values[:, np.newaxis] # 计算masked后的分组求和,并累加到结果 sum_n2 = (df1_n1 * n1_mask).groupby(level='n2', axis=1).sum() result = result.add(sum_n2, fill_value=0)
方法三:避免NaN,用乘法替代where
原代码的where本质是将不满足条件的元素设为NaN,求和时NaN会被忽略。我们可以直接用df1 * mask(不满足条件的元素设为0),求和结果一致,但0的处理比NaN高效:
# 1. 构建与df1形状匹配的mask矩阵 mask_df2 = df2.loc[df1.index] == 2 # 按df1列的n1映射到mask的对应列 mask_list = [mask_df2[col[0]] for col in df1.columns] mask_matrix = pd.concat(mask_list, axis=1) mask_matrix.columns = df1.columns # 2. 直接相乘后分组求和 result = (df1 * mask_matrix).groupby(level='n2', axis=1).sum()
注:此方法通过列映射确保mask与df1列完全匹配,适合所有场景,相比原where操作内存占用更低、计算更快。
内容的提问来源于stack exchange,提问作者Lei Hao
相关产品推荐
相关产品推荐

