Pandas:删除DataFrame分组尾部连续flag=1的行,排查代码问题
问题:删除DataFrame分组尾部连续flag=1的行
我有一个Pandas DataFrame,需要删除每个employeeid分组中尾部连续出现flag=1的所有行。以下是我的实现代码,但无法得到预期输出,帮忙排查问题:
import pandas as pd # 示例DataFrame df = pd.DataFrame({ 'employeeid': [1, 1, 1, 2, 2, 3, 3, 3], 'date': ['2022-01-01', '2022-01-02', '2022-01-03', '2022-01-01', '2022-01-02', '2022-01-01', '2022-01-02', '2022-01-03'], 'flag': [0, 1, 1, 1, 1, 0, 0, 1] }) df['date'] = pd.to_datetime(df['date']) df.sort_values(by=['employeeid', 'date'], ascending=False, inplace=True) mask = df.groupby('employeeid')['flag'].transform(lambda x: x[::-1].cumsum().eq(len(x)) & (x.iloc[-1] == 1)).astype(bool) df = df[~mask]
预期输出:
employeeid date flag 0 1 2022-01-01 0 5 3 2022-01-01 0 6 3 2022-01-02 0
问题分析
原代码的核心逻辑错误:
- 排序方向错误:使用
ascending=False将每个分组的日期倒序,后续处理逻辑混乱,无法准确识别"尾部"(最新日期)的连续1 - Mask生成逻辑错误:
x[::-1].cumsum().eq(len(x))仅当分组内所有flag都是1时才会成立,无法定位尾部连续的1,导致无法正确筛选需要保留的行
修正方案
方法1:分组自定义函数(直观易读)
先按employeeid和date升序排序,确保每个分组内的记录按时间从早到晚排列,尾部即为最新记录;再对每个分组找到最后一个非1的位置,保留该位置之前的所有行:
import pandas as pd df = pd.DataFrame({ 'employeeid': [1, 1, 1, 2, 2, 3, 3, 3], 'date': ['2022-01-01', '2022-01-02', '2022-01-03', '2022-01-01', '2022-01-02', '2022-01-01', '2022-01-02', '2022-01-03'], 'flag': [0, 1, 1, 1, 1, 0, 0, 1] }) df['date'] = pd.to_datetime(df['date']) # 升序排序,保证分组内时间从早到晚 df.sort_values(by=['employeeid', 'date'], ascending=True, inplace=True) def filter_trailing_ones(group): # 找到分组内最后一个flag≠1的索引 last_valid_idx = group[group['flag'] != 1].index.max() if pd.isna(last_valid_idx): # 分组全是1,直接返回空 return pd.DataFrame() # 返回从开头到最后一个非1的所有行 return group.loc[:last_valid_idx] # 应用分组过滤 df_cleaned = df.groupby('employeeid', group_keys=False).apply(filter_trailing_ones) print(df_cleaned)
方法2:Transform生成Mask(高效简洁)
利用反转分组后的累计求和,标记出尾部连续的1,生成过滤mask:
import pandas as pd df = pd.DataFrame({ 'employeeid': [1, 1, 1, 2, 2, 3, 3, 3], 'date': ['2022-01-01', '2022-01-02', '2022-01-03', '2022-01-01', '2022-01-02', '2022-01-01', '2022-01-02', '2022-01-03'], 'flag': [0, 1, 1, 1, 1, 0, 0, 1] }) df['date'] = pd.to_datetime(df['date']) df.sort_values(by=['employeeid', 'date'], ascending=True, inplace=True) # 生成过滤mask:保留非尾部连续1的行 mask = df.groupby('employeeid')['flag'].transform( lambda x: # 反转分组后,累计求和(统计从尾部开始的连续1数量) x[::-1].cumsum() # 遇到0时,后续累计值不再变化,标记为0 .where(x[::-1].cumsum() == x[::-1].cumsum().shift().fillna(0), 0) # 反转回原顺序,等于0的行就是需要保留的 [::-1] == 0 ) df_cleaned = df[mask] print(df_cleaned)
输出结果(两种方法均得到预期结果):
employeeid date flag 0 1 2022-01-01 0 5 3 2022-01-01 0 6 3 2022-01-02 0
内容的提问来源于stack exchange,提问作者r ram
相关产品推荐
相关产品推荐

