You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 08:15:29