如何在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:
然后实现单Id推理函数: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%置信区间上界 ])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
相关产品推荐
相关产品推荐

