如何在PanderaPolars/PanderaPandas中实现带过滤条件的列校验?
复杂多列联合校验在Pandera中的实现(Polars优先)
用户问题
我正尝试使用Pandera库对数据集进行数据质量校验,数据集加载为Polars DataFrame,采用PanderaPolars定义校验规则。部分列校验需设置多列联合条件,例如:
- 当column B = "unit_price"时,校验column A不为空
- 校验column C = column D + column E
- 校验column F与column G的相关性大于0.5
针对包含BUY_SELL(字符串类型,取值为B或S)、PRODUCT_TYPE(字符串类型,取值为REV或XX)、QUANTITY(浮点类型)的数据集,我需实现如下校验:当PRODUCT_TYPE = "REV"且BUY_SELL = "B"时,校验QUANTITY > 10。
我尝试用Pandas编写了两个自定义校验(代码如下),但未达到预期效果;且在Polars中无法编写等效逻辑,似乎不支持lambda表达式。请问这类复杂校验是否可在Pandera中实现(优先Polars,若不行Pandas也可)?
import pandera as pa import pandas as pd import warnings df = pd.read_csv('file.csv') check_quantity_grouped = pa.Check( lambda g: g[(True, "B")] > 10, groupby=lambda df: ( df.assign(product_type_rev=lambda d: d["PRODUCT_TYPE"] == "REV") .groupby(["product_type_rev", "BUY_SELL"]) ) ,ignore_na=True,raise_warning=False, error="trade quantity, when product_type = REV and BUY_SELL = B, is less than 10" ) check_quantity_filter = pa.Check( lambda df: df[(df['PRODUCT_TYPE'] == 'REV') & (df['BUY_SELL'] == 'B')] > 10 ,ignore_na=True,raise_warning=False ) schema = pa.DataFrameSchema({ "QUANTITY": pa.Column(float, [check_quantity_filter, check_quantity_grouped], nullable=True) }) try: schema(df, lazy=True) except pa.errors.SchemaErrors as exc: filtered_df = df[df.index.isin(exc.failure_cases["index"])] failures = exc.failure_cases print(f"filtered df:\n{filtered_df}") print(failures)
解决方案
一、Polars 版本实现(优先推荐)
Pandera对Polars的支持需使用pandera.api.polars模块,多列联合校验可通过自定义校验函数结合Polars表达式语法实现,无需lambda,逻辑更清晰。
代码示例
import polars as pl import pandera as pa from pandera.api.polars import DataFrameSchema, Column, Check # 构造测试数据 df = pl.DataFrame({ "BUY_SELL": ["B", "S", "B", "B"], "PRODUCT_TYPE": ["REV", "REV", "REV", "XX"], "QUANTITY": [5.0, 8.0, 15.0, 3.0] }) # 定义核心校验逻辑:符合条件的行必须满足QUANTITY>10,其他行跳过校验 def check_quantity_condition(df: pl.DataFrame) -> pl.Series: mask = (df["PRODUCT_TYPE"] == "REV") & (df["BUY_SELL"] == "B") return pl.when(mask).then(df["QUANTITY"] > 10).otherwise(True) # 定义Schema schema = DataFrameSchema({ "BUY_SELL": Column(pl.Utf8, Check.isin(["B", "S"])), "PRODUCT_TYPE": Column(pl.Utf8, Check.isin(["REV", "XX"])), "QUANTITY": Column(pl.Float64, Check(check_quantity_condition, error="当PRODUCT_TYPE=REV且BUY_SELL=B时,QUANTITY必须大于10")) }) # 执行校验 try: validated_df = schema.validate(df, lazy=True) print("校验通过") except pa.errors.SchemaErrors as exc: print("校验失败,错误行:") print(exc.failure_cases)
说明
- 自定义函数接收Polars DataFrame,返回与原数据长度一致的布尔Series:符合条件的行执行QUANTITY校验,其他行直接返回
True。 lazy=True模式会收集所有错误后再抛出,便于一次性查看所有违规行。
二、Pandas 版本修正
你之前的代码存在两个核心问题:
check_quantity_filter返回的是过滤后的Series,无法与原DataFrame的行一一对应,导致无法定位错误行。check_quantity_grouped的groupby参数逻辑错误,Pandera的groupby校验需要传入分组键,而非直接返回groupby对象。
修正后的代码
import pandas as pd import pandera as pa from pandera import DataFrameSchema, Column, Check # 构造测试数据 df = pd.DataFrame({ "BUY_SELL": ["B", "S", "B", "B"], "PRODUCT_TYPE": ["REV", "REV", "REV", "XX"], "QUANTITY": [5.0, 8.0, 15.0, 3.0] }) # 正确的逐行校验逻辑 check_quantity = Check( lambda df: df.apply( lambda row: row["QUANTITY"] > 10 if (row["PRODUCT_TYPE"] == "REV" and row["BUY_SELL"] == "B") else True, axis=1 ), ignore_na=True, error="当PRODUCT_TYPE=REV且BUY_SELL=B时,QUANTITY必须大于10" ) # 定义Schema schema = DataFrameSchema({ "BUY_SELL": Column(str, Check.isin(["B", "S"])), "PRODUCT_TYPE": Column(str, Check.isin(["REV", "XX"])), "QUANTITY": Column(float, [check_quantity], nullable=True) }) # 执行校验 try: validated_df = schema.validate(df, lazy=True) print("校验通过") except pa.errors.SchemaErrors as exc: print("校验失败,错误行:") print(exc.failure_cases) filtered_df = df.loc[exc.failure_cases["index"]] print("过滤后的错误数据集:") print(filtered_df)
说明
- 使用
df.apply(axis=1)逐行判断,确保返回的布尔Series长度与原DataFrame一致,能精准定位错误行。 - 移除无效的groupby校验,因为需求是逐行校验,无需分组逻辑。
其他复杂场景的实现思路
针对你提到的其他多列校验需求,可复用类似逻辑:
- 当column B="unit_price"时,校验column A不为空(Polars版本)
def check_a_not_null(df: pl.DataFrame) -> pl.Series: mask = df["B"] == "unit_price" return pl.when(mask).then(df["A"].is_not_null()).otherwise(True) - 校验column C = column D + column E(Polars版本)
check_c_sum = Check(lambda df: df["C"] == df["D"] + df["E"], error="C必须等于D+E") - 校验column F与column G的相关性大于0.5(Polars版本)
def check_correlation(df: pl.DataFrame) -> bool: corr = df.select(pl.corr("F", "G")).item() return corr > 0.5 check_fg_corr = Check(check_correlation, error="F和G的相关性必须大于0.5")
内容的提问来源于stack exchange,提问作者Zodiark1991
相关产品推荐
相关产品推荐

