Python-Polars如何使用字符串列表过滤分类列?
更新:
df_cat.filter(pl.col('a_cat').is_in(['a', 'c']))目前在Polars中可正常运行。
问题背景
我有如下Polars DataFrame:
df_cat = pl.DataFrame( [ pl.Series("a_cat", ["c", "a", "b", "c", "b"], dtype=pl.Categorical), pl.Series("b_cat", ["F", "G", "E", "G", "G"], dtype=pl.Categorical) ]) print(df_cat)
输出结果:
shape: (5, 2) ┌───────┬───────┐ │ a_cat ┆ b_cat │ │ --- ┆ --- │ │ cat ┆ cat │ ╞═══════╪═══════╡ │ c ┆ F │ │ a ┆ G │ │ b ┆ E │ │ c ┆ G │ │ b ┆ G │ └───────┴───────┘
使用单个字符串过滤可正常运行:
print(df_cat.filter(pl.col('a_cat') == 'c'))
输出结果:
shape: (2, 2) ┌───────┬───────┐ │ a_cat ┆ b_cat │ │ --- ┆ --- │ │ cat ┆ cat │ ╞═══════╪═══════╡ │ c ┆ F │ │ c ┆ G │ └───────┴───────┘
但尝试用字符串列表过滤时出现错误:
print(df_cat.filter(pl.col('a_cat').is_in(['a', 'c'])))
错误信息:
--------------------------------------------------------------------------- ComputeError Traceback (most recent call last) d:\GitRepo\Test2\stockEMD3.ipynb Cell 9 in <cell line: 1>() ----> 1 print(df_cat.filter(pl.col('a_cat').is_in(['c']))) File c:\ProgramData\Anaconda3\envs\charm3.9\lib\site-packages\polars\internals\dataframe\frame.py:2185, in DataFrame.filter(self, predicate) 2181 if _NUMPY_AVAILABLE and isinstance(predicate, np.ndarray): 2182 predicate = pli.Series(predicate) 2184 return ( -> 2185 self.lazy() 2186 .filter(predicate) # type: ignore[arg-type] 2187 .collect(no_optimization=True, string_cache=False) 2188 ) File c:\ProgramData\Anaconda3\envs\charm3.9\lib\site-packages\polars\internals\lazyframe\frame.py:660, in LazyFrame.collect(self, type_coercion, predicate_pushdown, projection_pushdown, simplify_expression, string_cache, no_optimization, slice_pushdown) 650 projection_pushdown = False 652 ldf = self._ldf.optimization_toggle( 653 type_coercion, 654 predicate_pushdown, (...) 658 slice_pushdown, 659 ) -> 660 return pli.wrap_df(ldf.collect()) ComputeError: joins/or comparisons on categorical dtypes can only happen if they are created under the same global string cache
疑问与解答
1. 为什么单个字符串的==过滤能正常运行?
当使用== 'c'这种单个字符串比较时,Polars会即时将单个字符串转换为对应分类的内部编码,这个过程不需要依赖全局字符串缓存——因为只涉及单个值的映射,Polars可以直接在当前分类列的字典中查找该字符串对应的编码,完成比较。
而is_in方法处理列表时,Polars需要先把整个字符串列表转换为分类类型,这个过程如果没有全局字符串缓存,会生成一个新的分类字典,和原列的分类字典不匹配,导致编码无法对应,从而抛出错误。
2. 使用字符串列表过滤分类列的正确方法是什么?
有两种可靠的解决方式:
方式一:启用全局字符串缓存
在创建DataFrame和执行过滤操作前,开启全局字符串缓存,确保所有分类相关的操作共享同一个字符串字典:
import polars as pl # 开启全局字符串缓存 pl.enable_string_cache() # 创建分类列DataFrame df_cat = pl.DataFrame( [ pl.Series("a_cat", ["c", "a", "b", "c", "b"], dtype=pl.Categorical), pl.Series("b_cat", ["F", "G", "E", "G", "G"], dtype=pl.Categorical) ]) # 列表过滤正常执行 result = df_cat.filter(pl.col('a_cat').is_in(['a', 'c'])) print(result)
方式二:将列表转换为与原列匹配的分类类型
手动将过滤列表转换为和目标列相同的分类类型,确保编码一致:
# 获取原列的分类字典 cat_dtype = df_cat['a_cat'].dtype # 将过滤列表转为对应分类 filter_list = pl.Series(["a", "c"], dtype=cat_dtype) # 执行过滤 result = df_cat.filter(pl.col('a_cat').is_in(filter_list)) print(result)
内容的提问来源于stack exchange,提问作者Peascod
相关产品推荐
相关产品推荐

