You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.12 08:31:01