基于层级条件聚合Pandas DataFrame列:小样本子类别合并
问题概述
现有一个Pandas DataFrame,Category列包含层级结构(类比国家-州-邮编,格式如A,B,123),Samples列为样本数量。需要对每个样本量不足5的子类别执行贪婪合并:将其与最接近的类别合并,若合并后仍不足5则继续合并,最终得到符合规则的结果。同时需支持扩展场景:若存在第三列,需根据列类型自动处理合并(整数列求和,字符串列拼接)。
初始DataFrame如下:
| Category | Samples |
|---|---|
| A,B,123 | 6 |
| A,B,456 | 3 |
| A,B,789 | 1 |
| X,Y,123 | 18 |
| X,Y,456 | 7 |
| X,Y,789 | 2 |
| P,Q,123 | 1 |
| P,Q,456 | 2 |
| P,Q,789 | 2 |
| L,M,123 | 1 |
| L,M,456 | 3 |
| S,T,123 | 5 |
| S,T,456 | 5 |
| S,T,789 | 3 |
合并规则明确如下:
- A,B场景:子类别456、789样本量均不足5,合并后仍不足5,需进一步与123合并,最终
A,B,123对应样本量10; - X,Y场景:仅子类别789样本量不足5,将其与最接近5的456合并,得到
X,Y,456对应样本量9,X,Y,123保留; - P,Q场景:所有子类别样本量均不足5,逐步合并后得到样本量为5的类别(如
P,Q,456); - L,M场景:两个子类别合并后仍不足5,保留合并结果
L,M,456对应样本量4; - S,T场景:仅789样本量不足5,可与123或456合并,两种结果均可行。
解决方案
步骤1:拆分Category层级
先将Category列按逗号拆分为多级列,方便按上层分组处理:
import pandas as pd # 构造初始DataFrame df = pd.DataFrame({ 'Category': ['A,B,123', 'A,B,456', 'A,B,789', 'X,Y,123', 'X,Y,456', 'X,Y,789', 'P,Q,123', 'P,Q,456', 'P,Q,789', 'L,M,123', 'L,M,456', 'S,T,123', 'S,T,456', 'S,T,789'], 'Samples': [6,3,1,18,7,2,1,2,2,1,3,5,5,3] }) # 拆分Category为三级 df[['L1', 'L2', 'L3']] = df['Category'].str.split(',', expand=True)
步骤2:实现贪婪合并函数
编写函数处理单个上层分组(比如L1=L2=A,B)内的合并逻辑,支持多字段类型处理:
def greedy_merge(group, value_cols=['Samples'], str_cols=None): str_cols = str_cols or [] # 按样本量降序排序,优先处理小样本项 group = group.sort_values('Samples', ascending=False).reset_index(drop=True) while True: # 筛选样本量不足5的项 small_mask = group['Samples'] < 5 if not small_mask.any(): break # 取第一个最小的不足项 small_row = group[small_mask].iloc[0] # 确定合并目标:优先找样本量≥5的项,没有则找除自身外的其他项 candidates = group[~small_mask] if (~small_mask).any() else group[group.index != small_row.name] if candidates.empty: break # 选择与5差值最小的项作为合并目标 candidates['diff'] = abs(candidates['Samples'] - 5) target_idx = candidates['diff'].idxmin() # 合并数值列:求和 for col in value_cols: group.loc[target_idx, col] += small_row[col] # 合并字符串列:用逗号拼接 for col in str_cols: group.loc[target_idx, col] = f"{group.loc[target_idx, col]},{small_row[col]}" # 删除被合并的行 group = group.drop(small_row.name).reset_index(drop=True) # 重新拼接Category列 group['Category'] = group[['L1', 'L2', 'L3']].agg(','.join, axis=1) return group.drop(['L1', 'L2', 'L3'], axis=1)
步骤3:应用合并逻辑到每个分组
按L1和L2分组,对每个分组应用合并函数:
# 按上层两级分组处理 merged_df = df.groupby(['L1', 'L2'], group_keys=False).apply(greedy_merge).reset_index(drop=True) # 查看结果 print(merged_df.sort_values('Category'))
扩展场景处理示例
如果新增一个字符串列Labels,可直接指定字符串列参数实现拼接:
# 新增字符串列 df['Labels'] = ['a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j', 'k', 'l', 'm', 'n'] # 调用合并函数时指定字符串列 merged_ext_df = df.groupby(['L1', 'L2'], group_keys=False).apply( lambda x: greedy_merge(x, str_cols=['Labels']) ).reset_index(drop=True)
结果验证
运行代码后,各场景结果符合要求:
- A,B分组:仅保留
A,B,123,Samples=10; - X,Y分组:保留
X,Y,123(18)和X,Y,456(9); - P,Q分组:保留
P,Q,456(5); - L,M分组:保留
L,M,456(4); - S,T分组:保留
S,T,123(8)和S,T,456(5) 或S,T,123(5)和S,T,456(8)(取决于合并目标选择)。
内容的提问来源于stack exchange,提问作者trojan horse
相关产品推荐
相关产品推荐

