Polars v1.32中map_batches的数组类型return_dtype参数配置问题
解决Polars v1.32+ map_batches指定return_dtype数组类型无效的问题
问题核心是你只指定了数组的元素类型pl.Int8,但没明确数组的固定形状——Polars的Array类型需要同时定义元素类型和形状,才能匹配你之前生成的Array(Int8, shape=(9,)) schema。
正确的return_dtype写法
直接用pl.Array(pl.Int8, shape=9)或者pl.Array(pl.Int8, shape=(9,))来指定完整的数组类型,而不是pl.Array[pl.Int8](这种写法只定义了元素类型,缺少形状信息)。
代码示例
假设你的原代码(v1.32前)是这样的:
import polars as pl import numpy as np def my_udf(batch: pl.DataFrame) -> np.ndarray: # 生成形状为(N,9)的int8数组,N是批次行数 return np.random.randint(-128, 127, size=(len(batch),9), dtype=np.int8) # v1.32前不需要指定return_dtype df = df.with_columns( pl.col("group_col").map_batches(my_udf).over("group_col").alias("signals") )
升级到v1.32+后,修改为:
df = df.with_columns( pl.col("group_col").map_batches( my_udf, return_dtype=pl.Array(pl.Int8, shape=(9,)) # 明确指定形状 ).over("group_col").alias("signals") )
额外说明
Polars v1.32开始强制要求map_batches指定return_dtype,是因为对于数组、结构体这类复杂类型,自动推断容易出错。你的UDF返回固定形状的numpy数组,必须明确告知Polars数组的形状,才能和之前的schema完全匹配。
内容的提问来源于stack exchange,提问作者Andi
相关产品推荐
相关产品推荐

