基于非相邻列差值过滤Pandas DataFrame行的实现方法
问题描述
现有如下自定义DataFrame对象df:
name val_1 val_2 val_3 val_4 AAA 1 2 3 11 BBB 2 3 5 9 CCC 6 4 15 10
需求为仅保留满足如下条件的行对应的name值:任意右侧val列相较其位置之前的任意val列数值增幅达到10及以上,不满足该条件的行予以删除。
已知diff()、ge()方法可用于差值比较,但差值计算不局限于相邻列时,不确定如何通过上述方法实现判断逻辑。
期望输出结果:
name AAA #val_4 increases by 10 from val_1 CCC #val_3 increases by 11 from val_2
最优实现方案
核心思路是利用numpy广播机制做批量向量化运算,避免Python层循环,性能最优,代码也足够简洁。
首先先构造测试数据方便复现:
import pandas as pd import numpy as np df = pd.DataFrame({ 'name': ['AAA', 'BBB', 'CCC'], 'val_1': [1, 2, 6], 'val_2': [2, 3, 4], 'val_3': [3, 5, 15], 'val_4': [11, 9, 10] }).set_index('name')
筛选符合条件的行
直接通过维度扩展批量计算所有列对的差值,一次性判断是否存在满足增幅≥10的组合:
# 提取所有val开头的数值列 val_cols = df.filter(like='val_') # 维度扩展后批量计算所有列两两之间的差值,判断是否存在≥10的增幅 # val_cols.values[:, :, None] 形状为(行数, 列数, 1) # val_cols.values[:, None, :] 形状为(行数, 1, 列数) # 相减后得到形状为(行数, 列数, 列数)的矩阵,存储每一行任意两列的差值 has_valid_diff = (val_cols.values[:, :, None] - val_cols.values[:, None, :] >= 10).any(axis=(1, 2)) # 筛选得到结果 result = df[has_valid_diff].reset_index()[['name']]
运行后result输出为:
name 0 AAA 1 CCC
补充:匹配差值说明
如果需要输出类似期望结果里的差值来源注释,可以加一段轻量遍历提取对应列对(仅对筛选后的行做遍历,数据量极小,几乎不影响性能):
col_list = val_cols.columns.tolist() for name in result['name']: row_data = val_cols.loc[name].values for i in range(len(col_list)): for j in range(i+1, len(col_list)): delta = row_data[j] - row_data[i] if delta >= 10: print(f"{name} #{col_list[j]} increases by {delta} from {col_list[i]}")
运行输出和预期完全一致:
AAA #val_4 increases by 10 from val_1 CCC #val_3 increases by 11 from val_2
关于diff()方法的说明
diff()默认仅计算相邻列的差值,若硬要用它实现需求,需要循环遍历不同的步长参数,反复做差再合并判断结果,代码冗余且性能远低于上面的广播方案,不推荐使用。
内容的提问来源于stack exchange,提问作者Roy
相关产品推荐
相关产品推荐

