如何编写函数移除pandas DataFrame中连续NaN超过阈值N的列
实现思路
- 逐列计算最大连续NaN长度,仅保留最大连续NaN长度≤阈值的列
- 核心逻辑通过布尔序列分组实现连续NaN的长度统计,计算效率高,适配千列级的DataFrame场景
代码实现
import pandas as pd def remove_consecutive_nan(df, threshold): def get_max_consecutive_nan(series): # 标记每个位置是否为NaN is_nan = series.isna() # 给连续的相同状态(NaN/非NaN)分配相同分组id group_id = is_nan.ne(is_nan.shift()).cumsum() # 计算所有NaN分组的长度,取最大值 nan_group_len = is_nan.groupby(group_id).sum() return nan_group_len.max() if len(nan_group_len) > 0 else 0 # 筛选符合要求的列 valid_cols = [col for col in df.columns if get_max_consecutive_nan(df[col]) <= threshold] return df[valid_cols]
示例测试
# 构造示例数据 data = { "ColA": [pd.NA] * 6, "ColB": [5, 6, 7, 5, 5, 3], "ColC": [3, pd.NA, 4, 5, 4, 3], "ColD": [pd.NA, 4, 4, pd.NA, pd.NA, pd.NA], "ColE": [pd.NA, 4, 4, pd.NA, 4, 3] } df = pd.DataFrame(data) threshold = 2 # 调用函数 res = remove_consecutive_nan(df, threshold) print(res)
运行后输出结果和你给出的期望结果完全一致。
内容的提问来源于stack exchange,提问作者kspr
相关产品推荐
相关产品推荐

