使用Polars拼接DataFrame后调用replace遇StringCacheMismatchError问题
Polars拼接多DataFrame后实现类别标签编码的解决方法
问题场景
当使用Polars拼接多个带有Categorical类型列的DataFrame后,尝试用replace方法将类别替换为整数时,会触发StringCacheMismatchError,报错信息如下:
StringCacheMismatchError: cannot compare categoricals coming from different sources, consider setting a global StringCache.
报错复现代码:
with pl.StringCache(): df1 = pl.DataFrame( {'a':[1,2], 'b':['a','b']}, schema = {'a':pl.Float32, 'b': pl.Categorical}) df2 = pl.DataFrame( {'c':[3,4], 'b':['a','c']}, schema = {'c': pl.Int32, 'b': pl.Categorical}) df = pl.concat([df1, df2], how = 'diagonal') # 执行replace操作时报错 cats = df['b'].cat.get_categories().to_list() df = df.with_columns( pl.col('b').replace(cats, range(len(cats)), return_dtype = pl.Int32))
报错原因
两个原始DataFrame的b列虽然都是Categorical类型,但它们的分类元数据来自不同的StringCache上下文,拼接后分类列的底层缓存不统一,replace操作时无法跨源比较分类值,导致报错。
解决方案
方案一:统一分类列上下文后使用replace
先将拼接后的分类列重新统一为同一组分类集合,确保缓存上下文一致,再执行replace操作:
# 统一b列的分类上下文 df = df.with_columns( pl.col('b').cat.set_categories(df['b'].cat.get_categories()) ) # 执行replace实现标签编码 cats = df['b'].cat.get_categories().to_list() df = df.with_columns( pl.col('b').replace(cats, range(len(cats)), return_dtype=pl.Int32) ) print(df)
方案二:直接使用Polars内置的分类编码(推荐)
Polars的Categorical列自带编码属性,通过cat.codes()可以直接获取对应的整数编码,无需手动用replace,既简洁又避免缓存问题:
df = df.with_columns( pl.col('b').cat.codes().cast(pl.Int32).alias('b') ) print(df)
执行后输出结果:
shape: (4, 3) ┌──────┬─────┬──────┐ │ a ┆ b ┆ c │ │ --- ┆ --- ┆ --- │ │ f32 ┆ i32 ┆ i32 │ ╞══════╪═════╪══════╡ │ 1.0 ┆ 0 ┆ null │ │ 2.0 ┆ 1 ┆ null │ │ null ┆ 0 ┆ 3 │ │ null ┆ 2 ┆ 4 │ └──────┴─────┴──────┘
内容的提问来源于stack exchange,提问作者figs_and_nuts
相关产品推荐
相关产品推荐

