在Pandas中创建布尔列,判断后续3行是否超过阈值12.3
问题:为Pandas DataFrame添加布尔列判断后续3行是否存在超过12.3的值
原始DataFrame
import pandas as pd df = pd.DataFrame({ 'other_stuff': ['lorem', 'ipsum', 'dolor', 'sit', 'amet', 'consectetur', 'adipiscing', 'elit', 'sed', 'do'], 'value': [12.0, 12.1, 11.9, 12.1, 12.4, 12.1, 12.2, 12.1, 11.8, 12.5] })
原始数据输出:
other_stuff value 0 lorem 12.0 1 ipsum 12.1 2 dolor 11.9 3 sit 12.1 4 amet 12.4 5 consectetur 12.1 6 adipiscing 12.2 7 elit 12.1 8 sed 11.8 9 do 12.5
需求说明
新增一个布尔列,规则如下:
- 当当前行后续3行(即索引
i的下一行i+1到i+3)的value列中存在数值超过12.3时,该列对应值为True - 若后续不足3行,则检查所有存在的后续行
- 无后续行时(如最后一行),值为
False
示例:
- 索引0的布尔值为
False,因为索引1、2、3的value均未超过12.3 - 索引1的布尔值为
True,因为索引2、3、4中存在超过12.3的数值(索引4的12.4)
期望最终结果
other_stuff value value > 12.3 in next 3 rows 0 lorem 12.0 False 1 ipsum 12.1 True 2 dolor 11.9 True 3 sit 12.1 True 4 amet 12.4 False 5 consectetur 12.1 False 6 adipiscing 12.2 True 7 elit 12.1 True 8 sed 11.8 True 9 do 12.5 False
解决方案
方法一:直观的shift逻辑判断
通过shift()获取后续行的value,再用逻辑或判断是否存在符合条件的值,最后填充缺失值:
# 分别获取后续1、2、3行的value,判断是否大于12.3,再取逻辑或 df['value > 12.3 in next 3 rows'] = ( (df['value'].shift(-1) > 12.3) | (df['value'].shift(-2) > 12.3) | (df['value'].shift(-3) > 12.3) ) # 最后几行的shift结果为NaN,替换为False df['value > 12.3 in next 3 rows'] = df['value > 12.3 in next 3 rows'].fillna(False)
方法二:高效的滚动窗口法(适合大规模数据)
先标记符合条件的行,再从后往前滚动窗口检查,最后调整位置得到结果:
# 先标记所有value>12.3的行 mask = df['value'] > 12.3 # 从后往前滚动窗口检查,再反转回来并偏移,得到当前行后续3行的结果 df['value > 12.3 in next 3 rows'] = mask[::-1].rolling(window=3, min_periods=1).any()[::-1].shift(1).fillna(False)
内容的提问来源于stack exchange,提问作者Jason Jarosz
相关产品推荐
相关产品推荐

