如何在Polars中实现分类模型预测结果的自定义条件排序?
二元分类样本的Polars自定义排序实现
需求说明
执行二元分类任务时,需要按以下规则查看样本:
- 优先级1:错误预测样本,按置信度从高到低排序(置信度最高的错误样本排最前)
- 优先级2:正确预测样本,按置信度从低到高排序(置信度最低的正确样本紧随错误样本之后)
目的是通过这类排序快速定位模型需要优化的样本类型(比如Stable Diffusion生成图像的分类场景),进而生成针对性训练样本。
数据示例
import polars as pl df = pl.from_repr(""" ┌──────────┬───────┬────────────┬────────────┬────────────────────┐ │ name ┆ truth ┆ prediction ┆ confidence ┆ correct_prediction │ │ --- ┆ --- ┆ --- ┆ --- ┆ --- │ │ str ┆ i64 ┆ i64 ┆ f64 ┆ i32 │ ╞══════════╪═══════╪════════════╪════════════╪════════════════════╡ │ Alice ┆ 1 ┆ 1 ┆ 0.343474 ┆ 1 │ │ Bob ┆ 0 ┆ 1 ┆ 0.298461 ┆ 0 │ │ Caroline ┆ 1 ┆ 1 ┆ 0.420634 ┆ 1 │ │ Dutch ┆ 0 ┆ 0 ┆ 0.125515 ┆ 1 │ │ Emily ┆ 1 ┆ 0 ┆ 0.772971 ┆ 0 │ │ Frank ┆ 0 ┆ 1 ┆ 0.646964 ┆ 0 │ │ Gerald ┆ 0 ┆ 0 ┆ 0.833705 ┆ 1 │ │ Henry ┆ 1 ┆ 1 ┆ 0.837181 ┆ 1 │ │ Isabelle ┆ 1 ┆ 1 ┆ 0.790773 ┆ 1 │ │ Jack ┆ 0 ┆ 0 ┆ 0.144983 ┆ 1 │ └──────────┴───────┴────────────┴────────────┴────────────────────┘ """)
期望排序结果:
expected = pl.from_repr(""" ┌──────────┬───────┬────────────┬────────────┬────────────────────┐ │ name ┆ truth ┆ prediction ┆ confidence ┆ correct_prediction │ │ --- ┆ --- ┆ --- ┆ --- ┆ --- │ │ str ┆ i64 ┆ i64 ┆ f64 ┆ i64 │ ╞══════════╪═══════╪════════════╪════════════╪════════════════════╡ │ Emily ┆ 1 ┆ 0 ┆ 0.772971 ┆ 0 │ │ Frank ┆ 0 ┆ 1 ┆ 0.646964 ┆ 0 │ │ Bob ┆ 0 ┆ 1 ┆ 0.298461 ┆ 0 │ │ Dutch ┆ 0 ┆ 0 ┆ 0.125515 ┆ 1 │ │ Jack ┆ 0 ┆ 0 ┆ 0.144983 ┆ 1 │ │ Alice ┆ 1 ┆ 1 ┆ 0.343474 ┆ 1 │ │ Caroline ┆ 1 ┆ 1 ┆ 0.420634 ┆ 1 │ │ Isabelle ┆ 1 ┆ 1 ┆ 0.790773 ┆ 1 │ │ Gerald ┆ 0 ┆ 0 ┆ 0.833705 ┆ 1 │ │ Henry ┆ 1 ┆ 1 ┆ 0.837181 ┆ 1 │ └──────────┴───────┴────────────┴────────────┴────────────────────┘ """)
解决方案
方法1:直接使用sort()中的条件表达式
Polars支持在sort()方法中传入多组排序规则,结合条件表达式实现优先级排序:
result = df.sort( # 第一优先级:错误样本排在前面(correct_prediction=0的样本权重更低,升序时排前) pl.col("correct_prediction"), # 第二优先级:错误样本按置信度降序,正确样本按置信度升序 pl.when(pl.col("correct_prediction") == 0) .then(pl.col("confidence") * -1) # 错误样本取负,升序等价于原置信度降序 .otherwise(pl.col("confidence")), ascending=[True, True] )
方法2:生成自定义排序列后排序
先计算一个辅助排序字段,再基于该字段排序:
result = df.with_columns( sort_key=pl.when(pl.col("correct_prediction") == 0) .then(pl.lit(0) - pl.col("confidence")) # 错误样本的排序键:0 - 置信度,值越小排越前 .otherwise(pl.lit(1) + pl.col("confidence")) # 正确样本的排序键:1 + 置信度,值越小排越前 ).sort("sort_key").drop("sort_key")
验证结果
两种方法都能得到与expected一致的排序结果:
print(result.equals(expected)) # 输出:True
说明
- 方法1更简洁,无需额外列,直接利用Polars的表达式能力实现多条件排序
- 方法2通过显式生成排序键,逻辑更直观,适合复杂排序规则的调试
内容的提问来源于stack exchange,提问作者Andrew Brēza
相关产品推荐
相关产品推荐

