Pandas:依据阈值对DataFrame进行动态层级分组
分层分组并按阈值终止后续分组的优化方案
需求说明
现有一个DataFrame,需按a、b、c依次执行分层分组,规则如下:
- 按当前层级分组后,若分组行数低于设定阈值,则该分组不再继续后续层级的分组
- 若分组行数高于阈值,则继续进行下一层级的分组
最终要得到示例中df_desired的结果,现有实现不够简洁,寻求更优方案。
示例数据
import pandas as pd import numpy as np df=pd.DataFrame({'a':[1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 2], 'b':[1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2], 'c':[1, 1, 1, 2, 2, 2, 3, 4, 4, 2, 2, 2, 2, 2, 2, 2, 2, 2] }) df_desired=pd.DataFrame({'a':[1, 1, 1, 2], 'b':[1, 1, 2, 2], 'c':[1, 2, np.nan, 2], 'group_size':[3, 3, 3, 9] })
现有实现(不够简洁)
group_cols = ['a', 'b', 'c'] grouped = list() n_threshold = 3 df_list = list() df_tmp = df.copy() for g in group_cols: print(g) grouped.append(g) df_grp = df_tmp.groupby(grouped).agg(n_rows=(g, 'size')).reset_index() df_lt = df_grp.loc[df_grp.n_rows <= n_threshold, ] df_list.append(df_lt) df_gt = df_grp.loc[df_grp.n_rows > n_threshold, ] df_tmp = pd.merge(df_tmp, df_gt[grouped], on = grouped, how = 'inner') df_list.append(df_gt) pd.concat(df_list)
优化实现方案
下面是更简洁高效的实现,核心思路是迭代处理每一层分组,分离终止分组和需继续分组的数据集,避免冗余的合并操作:
import pandas as pd import numpy as np def hierarchical_group(df, group_cols, threshold): result = [] # 初始待处理组:(数据集, 当前已用分组键列表) current_groups = [(df, [])] for col in group_cols: next_round = [] for data, used_keys in current_groups: # 按当前层级分组并计算大小 grouped_data = data.groupby(used_keys + [col]).size().reset_index(name='group_size') # 筛选出终止分组(行数≤阈值),补充后续未分组列为NaN stop_groups = grouped_data[grouped_data['group_size'] <= threshold] for unused_col in [c for c in group_cols if c not in used_keys + [col]]: stop_groups[unused_col] = np.nan result.append(stop_groups) # 筛选出需继续分组的数据集 continue_mask = grouped_data['group_size'] > threshold if continue_mask.any(): for _, row in grouped_data[continue_mask].iterrows(): # 生成过滤掩码,提取对应数据子集 filter_mask = True for k in used_keys + [col]: filter_mask &= (data[k] == row[k]) next_round.append((data[filter_mask], used_keys + [col])) current_groups = next_round if not current_groups: break # 无待处理组,提前结束 # 处理最后一轮剩余的待分组数据 for data, used_keys in current_groups: final_groups = data.groupby(used_keys).size().reset_index(name='group_size') for unused_col in [c for c in group_cols if c not in used_keys]: final_groups[unused_col] = np.nan result.append(final_groups) # 合并结果并整理列顺序 final_df = pd.concat(result, ignore_index=True) return final_df[group_cols + ['group_size']] # 测试运行 n_threshold = 3 df_result = hierarchical_group(df, ['a', 'b', 'c'], n_threshold) # 排序后匹配期望结果 df_result = df_result.sort_values(['a', 'b', 'c'], na_position='last').reset_index(drop=True) print(df_result)
优化点说明
- 去掉了原代码中重复的
merge操作,直接用掩码筛选数据,提升运行效率 - 逻辑更清晰:每一层明确分离终止和继续分组的部分,统一处理空值填充
- 支持任意长度的分组列列表,扩展性更强
- 加入提前终止逻辑,无待处理组时直接结束循环,减少无效计算
内容的提问来源于stack exchange,提问作者Helen
相关产品推荐
相关产品推荐

