如何基于Polars中另一DataFrame的断点对连续变量进行编码
如何基于Polars中另一DataFrame的断点对连续变量进行编码
咱们先来明确一下需求场景:我们有个存储断点规则的Polars DataFrame breakpoints,还有包含待编码连续特征的df,需要按照breakpoints里的规则把df的特征转换成整数编码,最终得到像encoded_df这样的结果。
首先看一下我们的断点数据:
import polars as pl breakpoints = pl.DataFrame( { "features": ["feature_0", "feature_0", "feature_1"], "breakpoints": [0.1, 0.5, 1], "n_possible_bins": [3, 3, 2], } ) print(breakpoints)
输出结果:
shape: (3, 3) ┌───────────┬─────────────┬─────────────────┐ │ features ┆ breakpoints ┆ n_possible_bins │ │ --- ┆ --- ┆ --- │ │ str ┆ f64 ┆ i64 │ ╞═══════════╪═════════════╪═════════════════╡ │ feature_0 ┆ 0.1 ┆ 3 │ │ feature_0 ┆ 0.5 ┆ 3 │ │ feature_1 ┆ 1.0 ┆ 2 │ └───────────┴─────────────┴─────────────────┘
然后是需要编码的原始数据df:
df = pl.DataFrame( {"feature_0": [0.05, 0.2, 0.6, 0.8], "feature_1": [0.5, 1.5, 1.0, 1.1]} ) print(df)
输出结果:
shape: (4, 2) ┌───────────┬───────────┐ │ feature_0 ┆ feature_1 │ │ --- ┆ --- │ │ f64 ┆ f64 │ ╞═══════════╪═══════════╡ │ 0.05 ┆ 0.5 │ │ 0.2 ┆ 1.5 │ │ 0.6 ┆ 1.0 │ │ 0.8 ┆ 1.1 │ └───────────┴───────────┘
我们期望得到的编码后结果encoded_df是这样的:
encoded_df = pl.DataFrame({"feature_0": [0, 1, 2, 2], "feature_1": [0, 1, 0, 1]}) print(encoded_df)
输出结果:
shape: (4, 2) ┌───────────┬───────────┐ │ feature_0 ┆ feature_1 │ │ --- ┆ --- │ │ i64 ┆ i64 │ ╞═══════════╪═══════════╡ │ 0 ┆ 0 │ │ 1 ┆ 1 │ │ 2 ┆ 0 │ │ 2 ┆ 1 │ └───────────┴───────────┘
还有几个关键规则要遵守:
- 可以保证待编码的所有特征都能在
breakpoints里找到对应规则 - 每个特征的标签是
np.array([str(i) for i in range(n_possible_bins)]),不同特征的分箱数可能不一样 - 编码要遵循
left_closed=False,也就是分箱区间是(断点, 下一个断点]的形式
现在问题来了:Polars的Expr.cut()方法虽然能实现分箱编码,但它的breaks参数需要传入一个浮点序列,怎么高效地从breakpoints这个DataFrame里提取每个特征对应的断点和标签来应用编码呢?
解决方案
我来给你捋个简单高效的实现思路:先把断点数据转换成一个便于查询的字典映射,再逐个特征应用cut()方法完成编码。
第一步:整理断点数据为特征映射
我们先把breakpoints按特征分组,整理出每个特征对应的断点列表和分箱数,转换成字典后后续用起来更方便:
import numpy as np # 分组聚合整理每个特征的参数 feature_params = ( breakpoints .group_by("features") .agg( pl.col("breakpoints").sort(), # 确保断点是升序的,cut需要有序断点 pl.col("n_possible_bins").first() ) .to_dict(as_series=False) ) # 转换成更易用的键值对结构 feature_map = {} for feat, breaks, n_bins in zip( feature_params["features"], feature_params["breakpoints"], feature_params["n_possible_bins"] ): feature_map[feat] = { "breaks": breaks, "labels": np.array([str(i) for i in range(n_bins)]), "n_bins": n_bins }
第二步:逐个特征应用编码
接下来遍历df的每个特征,用cut()方法按照对应的规则编码,最后转换成整数类型匹配目标结果:
# 构建编码表达式列表 encode_exprs = [] for feat in df.columns: params = feature_map[feat] # 构建cut表达式,设置left_closed=False,同时开启include_lowest确保最小值被正确分箱 encode_expr = ( pl.col(feat) .cut( breaks=params["breaks"], labels=params["labels"], left_closed=False, include_lowest=True ) .cast(pl.Int64) # 把字符串标签转成整数 .alias(feat) ) encode_exprs.append(encode_expr) # 生成编码后的DataFrame encoded_df = df.select(encode_exprs) print(encoded_df)
运行这段代码后,你就能得到和预期完全一致的编码结果啦。这里要注意两个细节:
include_lowest=True:因为我们设置了left_closed=False,开启这个参数能确保像feature_0里0.05这种小于第一个断点的值被正确分到第一个分箱(0号)- 一定要确保断点是升序排列的,不然
cut()会报错,所以我们在聚合的时候加了.sort()
备注:内容来源于stack exchange,提问作者Kevin Li
相关产品推荐
相关产品推荐

