Polars中对长度为1的列调用search_sorted结果变为标量的问题
问题分析与解决方案
问题重现
给定Polars DataFrame:
import polars as pl df = pl.DataFrame( [ pl.Series("ix", [0], dtype=pl.UInt32), pl.Series("s1", [[4, 8]], dtype=pl.List(pl.Int64)), pl.Series("s2", [[5]], dtype=pl.List(pl.Int64)), ] )
执行分组聚合的search_sorted操作:
res = df.group_by("ix").agg( pl.col("s1").explode().search_sorted(pl.col("s2").explode(), side="left").alias("r1"), pl.col("s2").explode().search_sorted(pl.col("s1").explode(), side="left").alias("r2"), )
得到结果中r1是标量(u32类型),而r2是列表(list[u32]类型),类型不一致。
原因分析
这并非bug,是Polars的类型推导逻辑导致的:
s2.explode()后是长度为1的单列数据,search_sorted接收单个元素时,返回的是标量值而非列表。- 聚合(
agg)操作会保留输入的类型:如果输入是标量,就直接存储为标量;如果输入是多个元素的序列,则自动包装为列表。而s1.explode()后是两个元素,search_sorted返回的是包含两个结果的序列,因此被聚合为列表。
修复方案
方案1:强制将标量转为列表
在聚合时直接用list.wrap()将结果包装为列表,确保类型统一:
res = df.group_by("ix").agg( pl.col("s1").explode().search_sorted(pl.col("s2").explode(), side="left").alias("r1").list.wrap(), pl.col("s2").explode().search_sorted(pl.col("s1").explode(), side="left").alias("r2"), )
此时r1会变为list[u32]类型,值为[1],和r2类型一致。
方案2:避免explode,直接在列表层面操作
使用list.eval在列表内部执行search_sorted,无需展开列表,天然保证结果为列表类型:
res = df.group_by("ix").agg( pl.col("s1").list.eval( pl.element().search_sorted(pl.col("s2").first(), side="left") ).alias("r1"), pl.col("s2").list.eval( pl.element().search_sorted(pl.col("s1").first(), side="left") ).alias("r2"), )
这种方式更贴合列表列的操作逻辑,避免了explode和重新聚合的开销,同时确保无论输入列表长度如何,输出都是列表类型。
验证结果
两种方案最终都会得到统一的类型:
shape: (1, 3) ┌─────┬───────┬───────────┐ │ ix ┆ r1 ┆ r2 │ │ --- ┆ --- ┆ --- │ │ u32 ┆ list[u32] ┆ list[u32] │ ╞═════╪═══════╪═══════════╡ │ 0 ┆ [1] ┆ [0, 1] │ └─────┴───────┴───────────┘
内容的提问来源于stack exchange,提问作者The Unfun Cat
相关产品推荐
相关产品推荐

