使用Prophet库分析多源数据时遭遇PyArrow类型错误求助
切换数据源时Prophet分析功能抛出ArrowTypeError错误
错误信息
PythonException: 从UDF抛出异常:
'pyarrow.lib.ArrowTypeError: 预期string或bytes类型,实际得到float64'
问题场景
我正在使用Prophet库开发数据分析功能,但切换不同数据源时遇到上述类型错误,核心代码如下:
def run_row_outlier_check(df: DataFrame, min_date, start_date, groupby_cols, job_id) -> DataFrame: pd_schema = StructType([ StructField(groupby_col, StringType(), True), StructField("ds", DateType(), True), StructField("y", IntegerType(), True), StructField("yhat", FloatType(), True), StructField("yhat_lower", FloatType(), True), StructField("yhat_upper", FloatType(), True), StructField("trend", FloatType(), True), StructField("trend_lower", FloatType(), True), StructField("trend_upper", FloatType(), True), StructField("additive_terms", FloatType(), True), StructField("additive_terms_lower", FloatType(), True), StructField("additive_terms_upper", FloatType(), True), StructField("weekly", FloatType(), True), StructField("weekly_lower", FloatType(), True), StructField("weekly_upper", FloatType(), True), StructField("yearly", FloatType(), True), StructField("yearly_lower", FloatType(), True), StructField("yearly_upper", FloatType(), True), StructField("multiplicative_terms", FloatType(), True), StructField("multiplicative_terms_lower", FloatType(), True), StructField("multiplicative_terms_upper", FloatType(), True) ]) # 生成连续日期的DataFrame df_rundates = (ps.DataFrame({'date':pd.date_range(start=min_date, end=(date.today() - timedelta(days=1)))})).to_spark() # 生成分组维度列表并关联日期,构建全量组合 df_bizlist = ( df.filter(f"{date_col} >= coalesce(date_sub(date 'today', {num_days_check}), '{start_date}')") .groupBy(groupby_col) .count() .orderBy(col("count").desc()) ) df_rundates_bus = ( df_rundates .join(df_bizlist, how='full') .select(df_bizlist[groupby_col], df_rundates["date"].alias("ds")) ) # 构建Prophet输入数据集 df_grouped_cnt = df.groupBy(groupby_cols).count() df_input = ( df_rundates_bus.selectExpr(f"{groupby_col}", "to_date(ds) as ds") .join(df_grouped_cnt.selectExpr(f"{groupby_col}", f"{date_col} as ds", "count as y"), on=['ds',f'{groupby_col}'], how='left') .withColumn("y", coalesce("y", lit(0))) .repartition(sc.defaultParallelism, "ds") ) # 缓存数据提升性能(注释状态) #df_input.cache().repartition(sc.defaultParallelism, "ds") # 按分组维度执行Prophet预测 df_forecast = ( df_input .groupBy(groupby_col) .applyInPandas(pd_apply_forecast, schema=pd_schema) ) # 过滤异常值并计算扣分 df_rowoutliers = ( df_forecast .filter("y > 0 AND (y > yhat_upper OR y < array_max(array(yhat_lower,0)))") .withColumn("check_type", lit("row_count")) .withColumn("deduct_score", expr("round(sqrt(pow(y-yhat, 2) / pow(yhat_lower - yhat_upper,2)))").cast('int')) .select( col("check_type"), col("ds").alias("ref_date"), col(groupby_col).alias("ref_dimension"), col("y").cast('int').alias("actual"), col("deduct_score"), col("yhat").alias("forecast"), col("yhat_lower").alias("forecast_lower"), col("yhat_upper").alias("forecast_upper") ) ) return add_metadata_columns(df_forecast, job_id), add_metadata_columns(df_rowoutliers, job_id) def add_metadata_columns(df: DataFrame, job_id) -> DataFrame: """ 为DataFrame添加任务元数据 """ df = df.select( lit(f"{job_id}").cast("int").alias("job_id"), lit(f"{source_type}").alias("source_type"), lit(f"{server}").alias("server_name"), lit(f"{database}").alias("database_name"), lit(f"{table}").alias("table_name"), lit(f"{groupby_col}").alias("ref_dim_column"), lit(f"{date_col}").alias("ref_date_column"), df["*"], expr("current_timestamp()").alias("_job_timestamp"), expr("current_user()").alias("_job_user") ).withColumnRenamed(groupby_col, "ref_dimension") return df
问题原因
错误源于Spark与Pandas通过PyArrow交互时的类型不匹配:
- 你定义的
pd_schema中,groupby_col被指定为StringType() - 但新数据源中该列实际是
float64类型(可能是数据源含空值自动转浮点,或列本身为数值型),导致PyArrow在序列化/反序列化时类型校验失败。
修复方案
1. 强制统一分组列类型
在数据预处理阶段,将groupby_col强制转为字符串类型,确保与Schema定义一致:
# 修改df_bizlist的分组列处理 df_bizlist = ( df.filter(f"{date_col} >= coalesce(date_sub(date 'today', {num_days_check}), '{start_date}')") .withColumn(groupby_col, col(groupby_col).cast(StringType())) # 强制转换为字符串 .groupBy(groupby_col) .count() .orderBy(col("count").desc()) ) # 修改df_input的join逻辑,确保两边分组列类型一致 df_input = ( df_rundates_bus.selectExpr(f"cast({groupby_col} as string) as {groupby_col}", "to_date(ds) as ds") .join(df_grouped_cnt.selectExpr(f"cast({groupby_col} as string) as {groupby_col}", f"{date_col} as ds", "count as y"), on=['ds', f'{groupby_col}'], how='left') .withColumn("y", coalesce("y", lit(0))) .repartition(sc.defaultParallelism, "ds") )
2. 校验数据源列类型
确认新数据源中groupby_col的原始类型,如果是数值型(如ID),转字符串不影响分组逻辑,但能彻底避免PyArrow类型冲突。
3. 验证Schema一致性
确保pd_schema中定义的所有列类型,与applyInPandas接收/返回的DataFrame列类型完全匹配,重点检查分组列、日期列等核心字段。
额外优化建议
add_metadata_columns函数中,lit(f"{groupby_col}")这类字符串拼接无需格式化,直接用lit(groupby_col)即可,减少潜在风险。- 根据数据量大小,启用
df_input.cache()可提升重复计算性能,建议开启。
内容的提问来源于stack exchange,提问作者Developer Rajinikanth
相关产品推荐
相关产品推荐

