You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何修改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

核心修改说明

  1. null值处理:先标记原序列的null位置,仅对非null子序列执行pl.cut,最后将结果回填到全量DataFrame中,null位置保持为None;
  2. 顺序维护:调整排序索引的生成逻辑,确保非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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.03 12:40:43