如何修改Polars自定义cut函数以正确处理Null值?
Polars自定义cut函数:同时解决顺序保留与null值处理问题
目前Polars的pl.cut函数存在两个明显问题:
- 无法处理包含null值的序列,直接执行会报错;
- 输出结果会按区间排序,无法保留原序列的元素顺序。
问题示例
1. 无法处理null值
执行以下代码会直接失败:
import polars as pl s = pl.Series([1, 1, 4, 3, 5, 2, 2, None]) pl.cut(s, bins=[2, 4])
2. 无法保留原序列顺序
执行代码:
import polars as pl s = pl.Series([1, 1, 4, 3, 5, 2, 2]) pl.cut(s, bins=[2, 4])
输出结果会打乱原序列顺序:
┌─────┬─────────────┬─────────────┐ │ ┆ break_point ┆ category │ │ --- ┆ --- ┆ --- │ │ f64 ┆ f64 ┆ cat │ ╞═════╪═════════════╪═════════════╡ │ 1.0 ┆ 2.0 ┆ (-inf, 2.0] │ │ 1.0 ┆ 2.0 ┆ (-inf, 2.0] │ │ 2.0 ┆ 2.0 ┆ (-inf, 2.0] │ │ 2.0 ┆ 2.0 ┆ (-inf, 2.0] │ │ 3.0 ┆ 4.0 ┆ (2.0, 4.0] │ │ 4.0 ┆ 4.0 ┆ (2.0, 4.0] │ │ 5.0 ┆ inf ┆ (4.0, inf] │ └─────┴─────────────┴─────────────┘
修改后的自定义cut函数
下面是修改后的函数,同时解决了顺序保留和null值处理问题——遇到null值时,对应的break_point和category字段直接返回null:
from typing import Optional import polars as pl def cut( s: pl.Series, bins: list[float], labels: Optional[list[str]] = None, break_point_label: str = "break_point", category_label: str = "category", maintain_order: bool = False, ) -> pl.DataFrame: # 记录原序列中null的位置 null_mask = s.is_null() # 提取非null子序列 non_null_s = s.filter(~null_mask) # 获取非null元素在原序列中的索引 non_null_indices = s.arg_true(~null_mask) if maintain_order: # 为非null序列生成排序索引,映射回原序列位置 _arg_sort_non_null = non_null_s.argsort() # 构建原序列全量排序索引 _arg_sort = pl.Series(name="_arg_sort", values=range(len(s))) _arg_sort = _arg_sort.set(non_null_indices, non_null_indices[_arg_sort_non_null]) # 仅对非null序列执行cut操作 result_non_null = pl.cut(non_null_s, bins, labels, break_point_label, category_label) if maintain_order: # 对非null结果按原顺序排序 result_non_null = ( result_non_null .select([pl.all(), pl.Series("_arg_sort", non_null_indices[_arg_sort_non_null])]) .sort('_arg_sort') .drop('_arg_sort') ) # 构建全量结果容器,初始值为null result = pl.DataFrame({ s.name: s, break_point_label: [None]*len(s), category_label: [None]*len(s) }) # 将非null部分的cut结果回填到对应位置 result = result.with_columns( pl.when(~null_mask) .then(result_non_null[break_point_label]) .otherwise(None) .alias(break_point_label), pl.when(~null_mask) .then(result_non_null[category_label]) .otherwise(None) .alias(category_label) ) return result
核心修改说明
- null值处理:先标记原序列的null位置,仅对非null子序列执行
pl.cut,最后将结果回填到全量DataFrame中,null位置保持为None; - 顺序维护:调整排序索引的生成逻辑,确保非null元素排序后能映射回原序列的正确位置,最终输出和原序列顺序完全一致。
测试验证
测试null值处理
s = pl.Series([1, 1, 4, 3, 5, 2, 2, None]) result = cut(s, bins=[2,4], maintain_order=True) print(result)
输出中最后一个元素的break_point和category均为null,且整体顺序和原序列一致。
测试顺序保留
s = pl.Series([1, 1, 4, 3, 5, 2, 2]) result = cut(s, bins=[2,4], maintain_order=True) print(result)
输出严格保留原序列顺序:
┌─────┬─────────────┬─────────────┐ │ ┆ break_point ┆ category │ │ --- ┆ --- ┆ --- │ │ f64 ┆ f64 ┆ cat │ ╞═════╪═════════════╪═════════════╡ │ 1.0 ┆ 2.0 ┆ (-inf, 2.0] │ │ 1.0 ┆ 2.0 ┆ (-inf, 2.0] │ │ 4.0 ┆ 4.0 ┆ (2.0, 4.0] │ │ 3.0 ┆ 4.0 ┆ (2.0, 4.0] │ │ 5.0 ┆ inf ┆ (4.0, inf] │ │ 2.0 ┆ 2.0 ┆ (-inf, 2.0] │ │ 2.0 ┆ 2.0 ┆ (-inf, 2.0] │ └─────┴─────────────┴─────────────┘
内容的提问来源于stack exchange,提问作者lebesgue
相关产品推荐
相关产品推荐

