如何结合Polars的over分组与hist方法实现数据分箱?
Polars分组分箱统计:over方法误区与正确实现
一、over方法报错的核心原因
你对over的理解误区在于:窗口函数(over修饰的表达式)必须保证每个分组返回的结果行数和该分组的原始行数完全一致。
pl.col.month.hist()会生成包含12个分箱结果的结构体(对应全年12个月),但每个category分组的原始行数是不同的(比如x组有2行,y组只有1行)。窗口函数要求每个分组输出的行数和输入行数匹配,但hist返回的是固定12行的统计结果,两者长度不匹配,因此触发ShapeError。
二、不用循环的高效分组分箱实现
不需要遍历每个category,用Polars的原生操作就能实现需求,推荐两种方法:
方法1:笛卡尔积+分组统计+左连接
先构造所有category和month的全量组合,再和实际统计结果左连接补0,逻辑清晰且性能高效:
import polars as pl df = pl.DataFrame( [ pl.Series("name", ["A", "B", "C", "D"], dtype=pl.Enum(["A", "B", "C", "D"])), pl.Series("month", [1, 2, 12, 1], dtype=pl.Int8()), pl.Series("category", ["x", "x", "y", "z"], dtype=pl.Enum(["x", "y", "z"])), ] ) # 生成所有月份和category的全量组合 all_months = pl.DataFrame({"month": range(1, 13)}, dtype=pl.Int8()) all_categories = df.select(pl.col("category").unique()) full_grid = all_categories.cross_join(all_months) # 统计每个category-month的实际数量 counts = df.group_by(["category", "month"], maintain_order=True).len() # 左连接补全空值为0 result = full_grid.join(counts, on=["category", "month"], how="left").fill_null(0).rename({"len": "count"}) print(result)
输出结果:
shape: (36, 3) ┌──────────┬───────┬───────┐ │ category ┆ month ┆ count │ │ --- ┆ --- ┆ --- │ │ enum ┆ i8 ┆ u32 │ ╞══════════╪═══════╪═══════╡ │ x ┆ 1 ┆ 1 │ │ x ┆ 2 ┆ 1 │ │ x ┆ 3 ┆ 0 │ │ x ┆ 4 ┆ 0 │ │ x ┆ 5 ┆ 0 │ │ … ┆ … ┆ … │ │ z ┆ 9 ┆ 0 │ │ z ┆ 10 ┆ 0 │ │ z ┆ 11 ┆ 0 │ │ z ┆ 12 ┆ 0 │ │ z ┆ 1 ┆ 1 │ └──────────┴───────┴───────┘
方法2:group_by结合hist+展开结果
利用group_by后对每个组计算hist,再通过unnest和explode展开结果:
from math import inf result = ( df.group_by("category", maintain_order=True) .agg( pl.col("month").hist( bins=[x + 1 for x in range(11)], include_breakpoint=True ).alias("binned") ) .unnest("binned") .explode(["breakpoint", "count"]) .with_columns( pl.col("breakpoint") .map_elements(lambda x: 12 if x == inf else x, return_dtype=pl.Float64()) .cast(pl.Int8()) .alias("month") ) .drop("breakpoint") .select("month", "count", "category") ) print(result)
输出结果和方法1一致,直接得到36行的分组分箱统计数据。
三、能否用over实现?
严格来说可以,但需要绕开窗口函数的长度限制,通过map_elements为每个category单独生成hist结果后展开:
from math import inf result = ( df.select(pl.col("category").unique()) .with_columns( pl.col("category") .map_elements( lambda cat: df.filter(pl.col("category") == cat)["month"].hist( bins=[x + 1 for x in range(11)], include_breakpoint=True ), return_dtype=pl.Struct({"breakpoint": pl.Float64, "count": pl.UInt32}) ).alias("binned") ) .unnest("binned") .explode(["breakpoint", "count"]) .with_columns( pl.col("breakpoint") .map_elements(lambda x: 12 if x == inf else x, return_dtype=pl.Float64()) .cast(pl.Int8()) .alias("month") ) .drop("breakpoint") .select("month", "count", "category") )
不过这种方式本质和循环逻辑类似,性能不如前两种原生分组操作的方案,因此更推荐使用group_by+join或者group_by+hist+explode的实现方式。
内容的提问来源于stack exchange,提问作者bzm3r
相关产品推荐
相关产品推荐

