从Pandas视角解决Polars的cut函数标签使用困惑
Pandas 迁移 Polars 时
cut 函数的标签逻辑优化方案 问题背景
我正在将代码从 Pandas 迁移至 Polars,使用 cut 函数时发现二者存在差异:
- Polars 无
bin参数,需自行计算断点; - 无法理解 Polars 中
label参数的使用逻辑,必须添加多余标签才能得到与 Pandas 一致的结果,想寻求更优实现方式。
原代码示例(Pandas 实现)
import numpy as np import pandas as pd import polars as pl # Polars DataFrame示例 data = { "value": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10], } df_pl = pl.DataFrame(data) # 转换为Pandas DataFrame以获取断点 df_pd = df_pl.to_pandas() # 使用retbins参数获取断点(来自Pandas) df_pd["cut_label_pd"], breakpoints = pd.cut(df_pd["value"], 4, labels=["low", "medium", "hight", "very high"], retbins=True) print(pl.from_pandas(df_pd))
Pandas 输出结果
shape: (10, 2) ┌───────┬──────────────┐ │ value ┆ cut_label_pd │ │ --- ┆ --- │ │ i64 ┆ cat │ ╞═══════╪══════════════╡ │ 1 ┆ low │ │ 2 ┆ low │ │ 3 ┆ low │ │ 4 ┆ medium │ │ 5 ┆ medium │ │ 6 ┆ hight │ │ 7 ┆ hight │ │ 8 ┆ very high │ │ 9 ┆ very high │ │ 10 ┆ very high │ └───────┴──────────────┘ print(breakpoints) # [ 0.991 3.25 5.5 7.75 10. ]
当前 Polars 实现(需添加多余标签)
# Polars中使用cut函数 labels = ["don't use it", "low", "medium", "hight", "very high", "don't use it too"] df_pl = df_pl.with_columns( pl.col("value").cut(breaks=breakpoints, labels=labels).alias("cut_label_pl") ) print(df_pl)
Polars 输出结果
shape: (10, 2) ┌───────┬──────────────┐ │ value ┆ cut_label_pl │ │ --- ┆ --- │ │ i64 ┆ cat │ ╞═══════╪══════════════╡ │ 1 ┆ low │ │ 2 ┆ low │ │ 3 ┆ low │ │ 4 ┆ medium │ │ 5 ┆ medium │ │ 6 ┆ hight │ │ 7 ┆ hight │ │ 8 ┆ very high │ │ 9 ┆ very high │ │ 10 ┆ very high │ └───────┴──────────────┘
优化解决方案
核心差异解析
Polars 和 Pandas 的 cut 标签逻辑本质不同:
- Pandas:断点数组长度为
n_bins + 1,标签数量等于n_bins,每个标签对应相邻两个断点的区间; - Polars:默认要求标签数量与断点数组长度一致,标签对应每个断点的「左侧区间」,同时用首尾标签处理超出断点范围的数值。
同时注意:Pandas 默认左闭右开区间,Polars 默认左开右闭,需通过 left_closed=True 对齐行为。
方案1:纯 Polars 实现(不依赖 Pandas 生成断点)
直接用 Polars 计算分位数得到断点,传入与区间数一致的标签即可:
import polars as pl data = {"value": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]} df_pl = pl.DataFrame(data) # 计算4分位数断点,对齐Pandas cut(4)逻辑 breakpoints = df_pl.select(pl.col("value").quantile([0, 0.25, 0.5, 0.75, 1])).to_numpy().flatten() # 微调第一个断点,避免最小值落在区间外(和Pandas逻辑一致) breakpoints[0] -= 1e-8 # 标签数量等于区间数(4个),直接传入 labels = ["low", "medium", "hight", "very high"] df_pl = df_pl.with_columns( pl.col("value").cut( breaks=breakpoints, labels=labels, left_closed=True # 对齐Pandas左闭右开的默认行为 ).alias("cut_label_pl") ) print(df_pl)
方案2:复用 Pandas 断点,简化标签映射
如果已经通过 Pandas 获取了断点,不需要添加多余标签,而是先获取区间再映射:
import pandas as pd import polars as pl data = {"value": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]} df_pl = pl.DataFrame(data) # 复用Pandas生成的断点 df_pd = df_pl.to_pandas() _, breakpoints = pd.cut(df_pd["value"], 4, labels=["low", "medium", "hight", "very high"], retbins=True) labels = ["low", "medium", "hight", "very high"] # 先cut得到区间字符串,再用map_dict映射为目标标签 df_pl = df_pl.with_columns( pl.col("value").cut(breaks=breakpoints, left_closed=True) .map_dict({f"[{breakpoints[i]:.3f}, {breakpoints[i+1]:.3f})": labels[i] for i in range(4)}) .alias("cut_label_pl") ) print(df_pl)
方案3:用 Polars qcut 替代(分位数分箱场景)
如果是基于分位数的分箱,直接用 Polars 的 qcut 函数,它支持直接指定分箱数和标签,无需手动计算断点:
import polars as pl data = {"value": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]} df_pl = pl.DataFrame(data) df_pl = df_pl.with_columns( pl.col("value").qcut( q=4, labels=["low", "medium", "hight", "very high"], left_closed=True ).alias("cut_label_pl") ) print(df_pl)
内容的提问来源于stack exchange,提问作者gillesa
相关产品推荐
相关产品推荐

