如何在Spark中实现特殊过滤:若分类存在N>3则保留该分类全部数据
解决方案
方法一:先筛选符合条件的分类再过滤
先提取所有存在N>3记录的Category,再用这些分类过滤原DataFrame:
import pandas as pd # 构造示例DataFrame data = { 'Cod': [1,1,1,1,1,1,1,1,2,3,3,3,3], 'Category': ['A','A','A','B','B','B','B','B','D','Z','Z','Z','Z'], 'N': [1,2,3,1,2,3,4,5,1,1,2,3,4] } df = pd.DataFrame(data) # 提取存在N>3的Category列表 valid_cats = df[df['N'] > 3]['Category'].unique() # 过滤得到结果 result = df[df['Category'].isin(valid_cats)] print(result)
输出结果:
Cod Category N 3 1 B 1 4 1 B 2 5 1 B 3 6 1 B 4 7 1 B 5 9 3 Z 1 10 3 Z 2 11 3 Z 3 12 3 Z 4
方法二:使用窗口函数(transform)标记过滤
通过groupby+transform给每个分组添加标记,判断该组是否存在N>3的记录,再过滤标记为True的行:
import pandas as pd import numpy as np # 构造示例DataFrame data = { 'Cod': [1,1,1,1,1,1,1,1,2,3,3,3,3], 'Category': ['A','A','A','B','B','B','B','B','D','Z','Z','Z','Z'], 'N': [1,2,3,1,2,3,4,5,1,1,2,3,4] } df = pd.DataFrame(data) # 给每行添加标记:所在Category是否有N>3的记录 df['valid_group'] = df.groupby('Category')['N'].transform(lambda x: np.any(x > 3)) # 过滤并删除标记列 result = df[df['valid_group']].drop('valid_group', axis=1) print(result)
输出结果和方法一完全一致。
内容的提问来源于stack exchange,提问作者OdiumPura
相关产品推荐
相关产品推荐

