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

如何在PySpark数据集上应用TensorFlow Probability STS模型实现分组预测

TFP STS模型在PySpark上按Id分组预测的最优实现方案

核心思路采用PySpark分组Pandas UDF + 分布式模型加载的方案,完全适配单Id时序独立预测的需求,可充分利用集群资源并行处理全量Id的预测任务,实现线性扩容。

实现步骤

  • 第一步:模型序列化存储
    先将训练好的STS模型通过tf.saved_model.save()或pickle序列化后,上传至HDFS、S3等所有Spark executor都可访问的分布式存储路径。如果你的STS模型有自定义预处理/后处理逻辑,建议封装为独立的推理函数和模型绑定存储,避免 executor 端依赖缺失。
  • 第二步:定义分组推理UDF
    采用PySpark的applyInPandas接口实现单Id分组推理,该接口会将每个Id的分组数据转换为pandas DataFrame输入,完美适配TFP模型的输入格式要求,参考代码如下:
    首先定义输出结果的Schema:
    from pyspark.sql.types import StructType, StructField, StringType, DateType, DoubleType
    
    output_schema = StructType([
        StructField("Id", StringType(), nullable=True),
        StructField("Date", DateType(), nullable=True),
        StructField("pred_value", DoubleType(), nullable=True),
        StructField("pred_lower", DoubleType(), nullable=True), # 95%置信区间下界
        StructField("pred_upper", DoubleType(), nullable=True) # 95%置信区间上界
    ])
    
    然后实现单Id推理函数:
    import pandas as pd
    
    def sts_inference_per_id(pdf: pd.DataFrame) -> pd.DataFrame:
        # pdf为单个Id对应的全量时序数据,已提前按Date排序
        import tensorflow as tf
        import tensorflow_probability as tfp
        sts = tfp.sts
    
        # 仅在executor首次调用函数时加载模型,避免重复加载开销
        if not hasattr(sts_inference_per_id, "trained_model"):
            sts_inference_per_id.trained_model = tf.saved_model.load("/distributed/path/to/your/sts_model")
        
        current_id = pdf["Id"].iloc[0]
        time_series = pdf["value"].values
        forecast_horizon = 30 # 替换为你实际需要的预测步长
    
        # 替换为你自己的STS推理逻辑
        forecast = sts_inference_per_id.trained_model.predict(
            time_series, 
            forecast_horizon=forecast_horizon
        )
        pred_mean = forecast.mean().numpy()
        pred_lower, pred_upper = forecast.distribution.quantile([0.025, 0.975]).numpy()
    
        # 生成预测对应的日期序列,替换为你实际的时间粒度
        pred_dates = pd.date_range(
            start=pdf["Date"].max(), 
            periods=forecast_horizon + 1, 
            freq="D" # 日粒度,可替换为W/M等你需要的粒度
        )[1:]
    
        # 构造返回结果
        return pd.DataFrame({
            "Id": [current_id] * forecast_horizon,
            "Date": pred_dates,
            "pred_value": pred_mean,
            "pred_lower": pred_lower,
            "pred_upper": pred_upper
        })
    
  • 第三步:触发全量推理
    对原始数据按Id分组后调用推理UDF即可得到全量Id的预测结果:
    # 先按Id、Date排序,保证每个Id的时序是正确的
    pred_result = original_df.orderBy("Id", "Date") \
                             .groupBy("Id") \
                             .applyInPandas(sts_inference_per_id, schema=output_schema)
    # 结果存储到目标路径
    pred_result.write.mode("overwrite").parquet("/your/result/path")
    

优化注意事项

  • 配置足够的executor内存:TF运行环境和模型加载需要一定内存,建议将executor内存设置为4G以上,避免出现OOM问题。
  • 控制批次大小:通过参数spark.sql.execution.arrow.maxRecordsPerBatch调整每个批次处理的Id数量,匹配模型的推理效率,避免 executor 负载过高。
  • 不需要额外做数据倾斜处理:你提到每个Id的记录条目数一致,不会出现分组数据倾斜的问题,并行效率可以达到最优。

内容的提问来源于stack exchange,提问作者THB

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 11:27:04