如何基于颜色阈值过滤Pandas多级索引DataFrame
多级索引DataFrame按不同colour设置阈值过滤数据
原DataFrame
| shape | colour | data | | | d1 | d2 | d3 | ------------------------------------------- | circle | green | 2 | 4 | 9 | | circle | red | -9 | 3 | 1 | | square | orange | 5 | -6 | 2 | | square | yellow | 9 | 8 | 2 |
通过以下代码创建:
import pandas as pd header = [ "shape", "colour", "label", "data" ] data = [ [ "circle", "green", "d1", 2 ], [ "circle", "green", "d2", 4 ], [ "circle", "green", "d3", 9 ], [ "square", "orange", "d1", 5 ], [ "square", "orange", "d2", -6 ], [ "square", "orange", "d3", 2 ], [ "circle", "red", "d1", -9 ], [ "circle", "red", "d2", 3 ], [ "circle", "red", "d3", 1 ], [ "square", "yellow", "d1", 9 ], [ "square", "yellow", "d2", 8 ], [ "square", "yellow", "d3", 2 ], ] raw = pd.DataFrame(data, columns=header) df = raw.pivot(index=["shape", "colour"], columns=["label"], values=["data"])
过滤规则
要求根据不同colour设置不同阈值,保留data列中绝对值大于对应阈值的值,不符合条件的设为NaN:
filter_rules = { "red": {"threshold": 5}, "green": {"threshold": 5}, "yellow": {"threshold": 7}, "orange": {"threshold": 1}, }
期望输出
| shape | colour | data | | | d1 | d2 | d3 | ------------------------------------------- | circle | green | Nan | Nan | 9 | | circle | red | -9 | Nan | Nan | | square | orange | 5 | -6 | 2 | | square | yellow | 9 | 8 | Nan |
尝试过的无效方法
以下两种方式均无法实现需求:
1.
df.apply(lambda r: r["data"] if abs(r["data"]) > filter[r.index.get_level_values(1)].get("threshold", 2) else 0)
df[(abs(df["data"]) > filter[df.index.get_level_values(1)].get("threshold", 0))]
正确解决方案
核心思路是生成与原DataFrame形状完全匹配的阈值矩阵,再逐元素进行比较过滤:
# 提取每行colour对应的阈值 thresholds = df.index.get_level_values("colour").map(lambda x: filter_rules[x]["threshold"]) # 将阈值扩展为与df列数一致的矩阵 threshold_matrix = pd.DataFrame([thresholds]*df.shape[1]).T # 对齐列名确保匹配 threshold_matrix.columns = df.columns # 使用where方法保留符合条件的值,不符合的替换为NaN result = df.where(abs(df) > threshold_matrix) print(result)
执行后输出结果与期望一致:
data label d1 d2 d3 shape colour circle green NaN NaN 9.0 red -9.0 NaN NaN square orange 5.0 -6.0 2.0 yellow 9.0 8.0 NaN
原理说明
map方法根据索引中的colour值匹配对应的阈值,得到每行的阈值序列- 将阈值序列重复扩展为与原DataFrame列数相同的矩阵,确保每个data列的元素都能匹配到对应colour的阈值
df.where()方法会保留满足条件(绝对值大于阈值)的原始值,不满足的替换为NaN
内容的提问来源于stack exchange,提问作者sam
相关产品推荐
相关产品推荐

