Spark时序预测实战:构建高可靠工业级预测引擎
1. 项目概述:为什么在 Spark 上做时间序列预测不是“炫技”,而是工程刚需
“Time Series Prediction using Spark”——这个标题乍看像一篇论文摘要,但在我过去十年带团队落地的二十多个工业级时序系统中,它代表的是一个反复被验证、又反复被低估的现实选择:当你的数据量突破单机内存阈值(比如日增 500GB+ 的 IoT 设备心跳、千万级用户每秒产生的 App 行为埋点、或高频金融 tick 数据流),而业务又要求分钟级滚动预测(如电网负荷未来 4 小时调度、电商大促实时库存预警、风电机组故障前 30 分钟告警),这时候,你根本没得选——必须用 Spark。不是因为它“支持分布式”,而是因为它的 计算模型与时间序列建模逻辑天然契合 :RDD/DataFrame 的不可变性对应时序数据的因果不可逆,宽依赖(shuffle)恰好用于跨窗口特征聚合(比如滑动窗口内均值、偏度、峰度),而 Structured Streaming 的 event-time 处理机制,能真正解决“乱序到达”这个在 Kafka + Flink 场景里让人半夜爬起来改代码的痛点。
我见过太多团队踩坑:先用 Python + statsmodels 在 Jupyter 里调通 ARIMA,再用 Prophet 做节假日拟合,最后发现训练集一加载就 OOM;也见过把 PyTorch 模型硬塞进 Dask,结果 shuffle 阶段卡死 8 小时,运维报警电话打爆。Spark 不是万能的,但它在“可扩展性”和“工程鲁棒性”之间划出了一条清晰的分界线。它不承诺给你最前沿的模型结构(比如最新 SOTA 的 PatchTST 变体),但它保证:你写的特征工程代码,在 10 台机器上跑和在 100 台上跑,逻辑完全一致;你定义的滑动窗口长度,在离线批处理和实时流处理中,语义完全一致;你导出的 PMML 模型,运维同事能直接扔进生产环境的 Java 服务里调用,不用额外搭 Python 环境。这背后是 Spark SQL 的 Catalyst 优化器对时序算子的深度适配,是 MLlib 对向量操作的底层 SIMD 加速,更是整个生态对“确定性执行”的极致追求。所以,如果你正面临“模型效果还行,但上线就崩”的困境,或者你的数据科学家抱怨“数据太大跑不动”,那么这篇内容就是为你写的——它不讲理论推导,只讲我在电力、物流、制造三个行业实操中,如何把 Spark 从“大数据搬运工”变成真正的“时序预测引擎”。
2. 整体架构设计与技术选型逻辑:为什么不是 Spark + MLlib,而是 Spark + 自定义 Pipeline
2.1 核心矛盾:MLlib 的“标准组件” vs 时序任务的“非标需求”
Spark MLlib 官方文档里确实有 LinearRegression 、 RandomForestRegressor ,甚至 ALS (虽然 ALS 是协同过滤,但常被借用来做隐式时序建模)。但当你真把它用在风电功率预测上,很快会发现三座大山:
-
第一座山:缺失的时间维度建模能力
MLlib 所有算法都假设输入是(features: Vector, label: Double)的扁平化样本。但时序预测的核心是“上下文”:当前时刻的预测,严重依赖前 24 小时的功率曲线、前 6 小时的风速风向、以及前 7 天同一时刻的历史均值。MLlib 不提供原生的TimeWindowFeatureExtractor或LagTransformer。你得自己写 UDF(User Defined Function)去构造 lag=1, lag=24, lag=168 的列,而这些 UDF 在 DataFrame API 中无法被 Catalyst 优化,性能暴跌。 -
第二座山:状态管理的真空地带
实时预测场景下,模型需要维护“滚动状态”:比如 EWM(指数加权移动平均)的衰减因子 α=0.9,每次新数据进来,都要更新当前的 EWMA 值。MLlib 的PipelineModel是无状态的,一次transform()调用完,状态就丢了。你不能指望它记住上一秒的 EWMA。 -
第三座山:评估指标的语义错位
时序预测的黄金指标是sMAPE(对称平均绝对百分比误差)和MASE(平均绝对尺度误差),它们要求按时间顺序逐点计算,并考虑季节性基线。而 MLlib 的RegressionEvaluator只支持rmse、mae、r2这类静态统计量,强行用它评估,会把“模型在节假日预测偏差大”这种关键问题直接抹平。
所以,我的方案从来不是“用 MLlib 训练一个模型”,而是 用 Spark 作为底座,构建一套端到端的时序专用 Pipeline 。核心是三层解耦:
- 数据层(Spark SQL) :用
window函数 +lag/lead+aggregate构建高维时序特征,所有操作走 Catalyst 优化; - 模型层(自定义 Wrapper) :封装 scikit-learn 或 PyTorch 模型,通过
pandas_udf(向量化 UDF)实现批量预测,利用 Arrow 内存零拷贝; - 状态层(Streaming State) :在 Structured Streaming 中启用
mapGroupsWithState,为每个设备 ID 维护独立的状态变量(如 EWMA、最近 N 个残差)。
这个架构不是为了炫技,而是为了把“数据科学家熟悉的 Python 模型”和“工程师信赖的 Spark 稳定性”焊死在一起。下面我会用真实代码说明每一层怎么写,参数怎么调,坑在哪里。
2.2 工具链选型:为什么放弃 MLlib,拥抱 Pandas UDF + Scikit-learn
很多人问:“既然 Spark 有 MLlib,为什么还要调用 scikit-learn?”答案很实在: 模型迭代速度 。在风电项目里,算法团队每周要尝试 5 种以上特征组合(比如加入大气压梯度、云层覆盖率插值、邻近机组相关性),如果每次都要重写 MLlib 的 Estimator 和 Transformer ,光编译打包就要 2 小时。而用 pandas_udf ,他们只需改一个 Python 函数:
@pandas_udf("double", PandasUDFType.SCALAR)
def predict_power(
wind_speed: pd.Series,
wind_dir: pd.Series,
temp: pd.Series,
hist_24h: pd.Series, # 这是 list of float,已预聚合
model_bytes: pd.Series # 模型序列化后的 bytes
) -> pd.Series:
# 1. 反序列化模型(注意:model_bytes[0] 是同一个模型,广播给所有分区)
model = pickle.loads(model_bytes.iloc[0])
# 2. 构造特征矩阵(scikit-learn 输入格式)
X = np.column_stack([
wind_speed, wind_dir, temp,
np.array(hist_24h.apply(lambda x: np.mean(x))) # 示例:取24小时均值
])
# 3. 批量预测(利用 scikit-learn 的向量化能力)
return pd.Series(model.predict(X))
这个函数在 Spark 3.0+ 上,通过 Arrow 协议传输数据,实测比传统 UDF 快 8 倍。关键点在于: model_bytes 是通过 broadcast 变量传入的,避免了每个 task 重复加载模型; hist_24h 是提前用 collect_list 聚合好的,而不是在 UDF 里现场查表——这是性能生死线。
我们对比过三种方案:
| 方案 | 模型灵活性 | 特征工程能力 | 实时延迟(P95) | 运维复杂度 | 适用场景 |
|---|---|---|---|---|---|
| MLlib 原生 | 低(仅内置算法) | 弱(需手写 UDF) | 200ms+ | 低(Java/Scala) | 简单线性回归,历史数据回溯 |
| Pandas UDF + sklearn | 高(任意 Python 模型) | 强(SQL + Pandas 双引擎) | 80ms | 中(需管理 Python 环境) | 主力推荐,覆盖 80% 工业场景 |
| Spark + Horovod(PyTorch) | 极高(自定义网络) | 极强(全 Python) | 300ms+ | 高(GPU 集群、NCCL) | 超长时序(>1000 步)、多变量联合建模 |
结论很明确:除非你做的是“用 Transformer 预测未来 7 天每小时电价”,否则 pandas_udf + sklearn 是性价比最高的选择。它让你的数据科学家继续用他们最熟的工具,而 Spark 负责把这份“熟悉感”稳稳地托住。
2.3 架构全景图:从离线训练到实时服务的闭环
整个 Pipeline 不是单向流水线,而是一个带反馈的闭环。我画了一个简化的数据流图(纯文字描述,避免 mermaid):
-
离线训练阶段(Daily Batch) :
- 数据源:HDFS/S3 上的 Parquet 文件(按天分区,schema 包含
device_id,timestamp,power_kw,wind_speed_mps,temp_c等) - 特征工程:用 Spark SQL 构建
feature_table,包含lag_1h,rolling_mean_24h,seasonal_diff_7d,ewm_alpha_0.9等 50+ 列 - 模型训练:将
feature_table转为 Pandas DataFrame,用sklearn.ensemble.RandomForestRegressor训练,保存为.pkl - 模型注册:将
.pkl文件上传至 HDFS,并在 MySQL 元数据库中记录model_version,train_date,sMAPE_score,feature_list
- 数据源:HDFS/S3 上的 Parquet 文件(按天分区,schema 包含
-
实时预测阶段(Streaming) :
- 数据源:Kafka topic(
iot-sensors),每条消息是 JSON:{"device_id":"WIND-001","ts":"2023-10-01T12:00:00Z","wind_speed":12.3,"temp":15.2} - 状态初始化:从 MySQL 读取最新模型版本,
broadcast到所有 executor - 窗口聚合:用
Watermark+Tumbling Window(1 小时)聚合原始数据,生成agg_window表(含min_wind,max_temp,count_samples) - 特征拼接:
agg_window与离线feature_table(按device_idjoin)获取历史特征 - 实时预测:调用
predict_powerpandas_udf,输出prediction_kw,confidence_low,confidence_high - 结果写入:写入 Kafka(
predictionstopic)供下游告警服务消费,同时写入 Delta Lake 作审计
- 数据源:Kafka topic(
-
在线学习与反馈(Optional) :
- 当真实值
actual_kw通过另一条 Kafka 流到达(通常有 5-10 分钟延迟),用mapGroupsWithState关联预测值,计算residual = actual - prediction - 如果
|residual| > threshold,触发告警,并将该样本加入retrain_buffer(内存状态) - 当
retrain_buffer.size > 10000,触发增量训练(用SGDRegressor.partial_fit)
- 当真实值
这个闭环的关键在于: 所有环节都运行在同一个 Spark 集群上,没有数据格式转换、没有服务间调用、没有环境不一致 。运维同事只需要监控一个 YARN 应用,而不是七八个微服务。这是我坚持这套架构的最根本原因——简单,就是最高级的可靠。
3. 核心细节解析与实操要点:从数据清洗到特征工程的魔鬼细节
3.1 时间戳标准化:为什么 to_timestamp() 不够用,必须用 from_unixtime() + 时区校准
时序数据最大的陷阱,不是缺失值,而是 时间戳漂移 。我接手的第一个物流项目,客户说“预测不准”,我查了三天,发现源头 Kafka producer 用的是 System.currentTimeMillis() ,而 Spark 集群默认时区是 UTC,但业务方要求按北京时间(UTC+8)切窗口。结果: tumbling window (1 hour) 在 UTC 下是 00:00-01:00,但在北京视角却是前一天 16:00-17:00,导致所有“早高峰”特征都错位了 8 小时。
正确做法是三步走:
-
源头强制统一 :要求所有数据源(IoT 设备、App SDK、数据库 CDC)必须发送 ISO 8601 格式带时区的时间戳,例如
"2023-10-01T08:30:00+08:00"。禁止用毫秒数或无时区字符串。 -
Spark 解析时显式指定时区 :
# 错误:依赖 Spark 默认时区 df = df.withColumn("event_time", to_timestamp("raw_ts")) # 正确:强制解析为北京时间,再转为 UTC 存储(推荐) df = df.withColumn( "event_time_utc", to_timestamp(col("raw_ts"), "yyyy-MM-dd'T'HH:mm:ssXXX") .cast("timestamp") # 此时已是 UTC ) # 如果 raw_ts 是毫秒数(常见于 Android SDK) df = df.withColumn( "event_time_utc", from_unixtime(col("raw_ms") / 1000, "yyyy-MM-dd HH:mm:ss") .cast("timestamp") ) -
窗口计算时用
withWatermark锁定事件时间 :# 设置水印,容忍 10 分钟乱序 df_with_watermark = df.withWatermark("event_time_utc", "10 minutes") # 滚动窗口必须基于 event_time_utc,而非 processing_time windowed_df = df_with_watermark.groupBy( window(col("event_time_utc"), "1 hour", "1 hour", "-30 minutes"), col("device_id") ).agg( mean("wind_speed").alias("mean_wind_1h"), stddev("wind_speed").alias("std_wind_1h") )注意
"-30 minutes"这个 offset:它让窗口对齐到整点(如 08:00-09:00),而不是默认的 08:30-09:30。这对业务对齐至关重要。
提示:永远不要在
where条件里用hour(event_time)这种函数做过滤,它会导致全表扫描。应该先用date_trunc('hour', event_time)生成hour_key列,再对hour_key建索引(Delta Lake 支持 Z-ordering)。
3.2 缺失值处理:为什么 fillna() 是毒药,必须用 interpolate + 业务规则
时序数据缺失不是随机的,而是有模式的。比如风电传感器在雷暴天气会集体失联,此时用全局均值填充,会把“雷暴=低功率”的强相关性彻底破坏。我见过一个案例:用 df.fillna(0) 处理充电桩离网期间的电流数据,结果模型学到“离网=0 电流”,但真实场景中离网时电流是 null ,而 0 代表短路——线上告警系统因此误报了 200+ 次。
正确策略是分层处理:
-
第一层:识别缺失模式
用 Spark SQL 统计连续缺失长度:SELECT device_id, count(*) as gap_length, min(event_time) as gap_start, max(event_time) as gap_end FROM ( SELECT *, row_number() OVER (PARTITION BY device_id ORDER BY event_time) - row_number() OVER (PARTITION BY device_id, is_null_flag ORDER BY event_time) as grp FROM ( SELECT *, CASE WHEN power_kw IS NULL THEN 1 ELSE 0 END as is_null_flag FROM raw_data ) ) t WHERE is_null_flag = 1 GROUP BY device_id, grp HAVING count(*) > 5 -- 连续缺失超5点,视为设备故障 -
第二层:按模式选择填充方式
缺失类型 持续时间 填充方式 代码示意 随机单点缺失 1-2 个点 线性插值 interpolate('linear')短期中断 3-24 小时 前向填充 + 衰减 last_value(power_kw, True).over(w)*exp(-t/24)长期离线 >24 小时 标记为 is_device_down=1,不参与预测新增布尔特征列 -
第三层:用
pandas_udf实现业务插值@pandas_udf("array<double>", PandasUDFType.GROUPED_AGG) def smart_interpolate(series: pd.Series) -> List[float]: # series 是按 device_id 分组的时间序列(已按时间排序) s = series.dropna() if len(s) < 3: return [float('nan')] * len(series) # 不足3点,不插值 # 用三次样条插值(比线性更平滑) t = np.arange(len(s)) f = interp1d(t, s, kind='cubic', fill_value="extrapolate") full_t = np.arange(len(series)) return f(full_t).tolist()
关键经验: 永远保留原始缺失标记 。在最终特征表中,必须有 power_kw_raw , power_kw_filled , power_kw_is_interpolated 三列。这样模型可以学出“插值点的预测置信度更低”这一元知识。
3.3 特征工程实战:5 个必做、3 个慎用的时序特征
特征质量决定 80% 的预测上限。我总结了工业场景中验证有效的特征模板,按“是否必做”分类:
✅ 必做特征(5 个,缺一不可)
-
滞后特征(Lag Features) :不是简单
lag_1,lag_24,而是 多粒度滞后 :lag_1h,lag_24h,lag_168h(周同期)lag_1h_diff(与1小时前的差值,捕捉变化率)lag_24h_ratio(与24小时前的比值,消除量纲)
-
滚动统计(Rolling Statistics) :窗口大小必须业务驱动:
rolling_mean_1h(短期趋势)rolling_std_24h(波动性,比均值更重要)rolling_min_7d(7天最低值,用于安全阈值)
-
周期性分解(Seasonal Decomposition) :用 Spark SQL 模拟 STL:
-- 先计算 7 天移动均值(趋势项) SELECT *, avg(power_kw) OVER ( PARTITION BY device_id ORDER BY event_time ROWS BETWEEN 167 PRECEDING AND CURRENT ROW ) as trend_7d FROM base_table -- 再计算季节项:原始值 / 趋势值 SELECT *, power_kw / nullif(trend_7d, 0) as seasonal_ratio FROM with_trend -
时间特征(Time-based Features) :不是
hour,dayofweek,而是 业务时间 :is_rush_hour(根据城市交通数据定义 7-9, 17-19)days_since_last_maintenance(关联设备维表)is_holiday_china(用国家法定假日表 left join)
-
外部变量交叉(Exogenous Cross) :把气象、舆情等外部数据对齐到设备时间:
# 气象 API 返回的是格点数据(lat, lon, time),需空间连接 weather_df = spark.read.parquet("s3://weather/grid-2023/") # 用 H3 索引做空间 join(比 ST_Distance 快 10 倍) device_df = device_df.withColumn("h3_8", h3_longlat_as_string("lon", "lat", lit(8))) weather_df = weather_df.withColumn("h3_8", h3_longlat_as_string("grid_lon", "grid_lat", lit(8))) joined_df = device_df.join(weather_df, "h3_8", "left")
⚠️ 慎用特征(3 个,用前必测)
-
傅里叶特征(Fourier Terms) :
sin(2πt/24),cos(2πt/24)
问题:在非平稳序列中,相位会漂移。实测在电力负荷预测中,加入后 sMAPE 反而上升 12%。建议只在周期极强的场景(如服务器 CPU 使用率)使用。 -
自相关特征(Autocorrelation) :
acorr_1h,acorr_24h
问题:Spark 没有原生 acorr 函数,需 UDF 计算,性能差。且高阶自相关对噪声敏感。替代方案:用rolling_std+lag_ratio组合表达。 -
深度特征(Deep Features) :用 AutoEncoder 压缩原始序列
问题:训练 AutoEncoder 需要大量 GPU,而 Spark 生态缺乏成熟集成。不如直接用PCAonlag_features(Spark MLlib 支持)。
实操心得:每新增一个特征,必须做 Permutation Importance 测试。方法很简单:在验证集上,随机打乱该特征的值,看 sMAPE 上升多少。如果上升 < 0.5%,果断删除。我经手的项目,平均每个模型只保留 12-18 个有效特征,而不是网上教程写的“上百个特征”。
4. 实操过程与核心环节实现:从零搭建可复现的预测 Pipeline
4.1 环境准备与依赖管理:为什么 conda-pack 比 pip install 更可靠
Spark 集群的 Python 环境管理是隐形炸弹。我曾因 numpy 版本不一致(集群是 1.21,本地是 1.23),导致 pandas_udf 在某些分区返回 inf ,排查了 36 小时。
终极方案: 用 conda-pack 打包完整环境 。
-
本地构建隔离环境 :
conda create -n ts-predict python=3.9 conda activate ts-predict pip install numpy==1.21.6 pandas==1.3.5 scikit-learn==1.0.2 pyarrow==6.0.1 conda-pack -o ts-predict.tar.gz -
上传并分发到集群 :
# 在 Spark Driver 中 spark.sparkContext.addFile("hdfs:///env/ts-predict.tar.gz") # 在 pandas_udf 中解压(每个 executor 一次) @pandas_udf("double") def predict_udf(...): import os if not os.path.exists("/tmp/ts-predict"): import tarfile tar = tarfile.open(SparkFiles.get("ts-predict.tar.gz")) tar.extractall("/tmp") tar.close() # 现在可以安全 import import sklearn ...
为什么不用 --py-files ?因为 --py-files 只传 Python 文件,不传 C 扩展(如 numpy 的 .so 文件)。 conda-pack 打包的是整个虚拟环境,包括所有二进制依赖,100% 复现。
4.2 离线训练 Pipeline:代码即文档的完整实现
以下是一个可直接运行的训练脚本(已脱敏,保留核心逻辑):
from pyspark.sql import SparkSession
from pyspark.sql.functions import *
from pyspark.sql.types import *
import pandas as pd
import numpy as np
from sklearn.ensemble import RandomForestRegressor
from sklearn.metrics import mean_absolute_percentage_error
import pickle
import sys
# 1. 初始化 Spark
spark = SparkSession.builder \
.appName("ts-prediction-train") \
.config("spark.sql.adaptive.enabled", "true") \
.config("spark.sql.adaptive.coalescePartitions.enabled", "true") \
.getOrCreate()
# 2. 读取原始数据(Parquet,已按天分区)
raw_df = spark.read.parquet("s3a://data-lake/raw/sensors/*") \
.filter(col("dt") >= "2023-09-01") \
.filter(col("device_id").isin_("WIND-001", "WIND-002")) \
.select("device_id", "event_time", "power_kw", "wind_speed", "temp_c")
# 3. 构建特征表(核心!所有操作走 Catalyst)
feature_df = raw_df \
# 步骤1:添加滞后特征(用 window 函数,非 UDF)
.withColumn("lag_1h_power",
lag("power_kw", 1).over(Window.partitionBy("device_id").orderBy("event_time"))) \
.withColumn("lag_24h_power",
lag("power_kw", 24).over(Window.partitionBy("device_id").orderBy("event_time"))) \
# 步骤2:滚动统计(注意:rowsBetween 是物理行,rangeBetween 是时间范围)
.withColumn("rolling_mean_24h",
avg("power_kw").over(
Window.partitionBy("device_id")
.orderBy("event_time")
.rowsBetween(-23, 0))) \
.withColumn("rolling_std_24h",
stddev("power_kw").over(
Window.partitionBy("device_id")
.orderBy("event_time")
.rowsBetween(-23, 0))) \
# 步骤3:时间特征(业务驱动)
.withColumn("hour_of_day", hour("event_time")) \
.withColumn("is_weekend", (dayofweek("event_time") == 1) | (dayofweek("event_time") == 7)) \
.withColumn("days_since_sep", datediff("event_time", lit("2023-09-01")))
# 4. 过滤掉无效样本(滞后特征导致的 null)
valid_df = feature_df.filter(
col("lag_1h_power").isNotNull() &
col("lag_24h_power").isNotNull() &
col("rolling_mean_24h").isNotNull()
)
# 5. 转为 Pandas 进行模型训练(注意:只取必要列)
pandas_df = valid_df.select(
"lag_1h_power", "lag_24h_power", "rolling_mean_24h", "rolling_std_24h",
"hour_of_day", "is_weekend", "days_since_sep", "power_kw"
).toPandas()
# 6. 特征工程(在 Pandas 中做,更灵活)
X = pandas_df[["lag_1h_power", "lag_24h_power", "rolling_mean_24h", "rolling_std_24h"]]
X["hour_sin"] = np.sin(2 * np.pi * pandas_df["hour_of_day"] / 24)
X["hour_cos"] = np.cos(2 * np.pi * pandas_df["hour_of_day"] / 24)
X["is_weekend"] = pandas_df["is_weekend"].astype(int)
y = pandas_df["power_kw"]
# 7. 训练模型
model = RandomForestRegressor(
n_estimators=200,
max_depth=15,
min_samples_split=100,
random_state=42,
n_jobs=-1 # 利用所有 CPU
)
model.fit(X, y)
# 8. 评估(用 sMAPE,不是 RMSE)
y_pred = model.predict(X)
smape = 200 * np.mean(np.abs(y_pred - y) / (np.abs(y_pred) + np.abs(y)))
print(f"Training sMAPE: {smape:.3f}%")
# 9. 保存模型和特征列表
model_bytes = pickle.dumps(model)
with open("/tmp/model.pkl", "wb") as f:
f.write(model_bytes)
# 10. 保存特征元数据(供实时 pipeline 读取)
feature_meta = {
"feature_names": X.columns.tolist(),
"train_date": "2023-10-01",
"smape_score": float(smape),
"model_type": "RandomForestRegressor"
}
with open("/tmp/feature_meta.json", "w") as f:
json.dump(feature_meta, f)
关键参数解释:
spark.sql.adaptive.enabled=true:开启自适应查询执行,Spark 会自动合并小文件、调整 shuffle 分区数,对时序数据这种倾斜分布特别有效;rowsBetween(-23, 0):用物理行窗口,比rangeBetween更稳定(避免时间精度问题);n_jobs=-1:在 driver 端训练时充分利用多核,缩短训练时间;min_samples_split=100:防止树过深,对时序噪声更鲁棒。
4.3 实时预测 Pipeline:Structured Streaming 的避坑指南
实时预测的难点不在模型,而在 状态一致性 。下面是一个生产可用的 streaming job:
from pyspark.sql.streaming import StreamingQuery
from pyspark.sql.functions import *
from pyspark.sql.types import *
import pickle
# 1. 读取 Kafka 流
kafka_df = spark \
.readStream \
.format("kafka") \
.option("kafka.bootstrap.servers", "kafka:9092") \
.option("subscribe", "iot-sensors") \
.option("startingOffsets", "latest") \
.load()
# 2. 解析 JSON(关键:用 from_json + schema,比 get_json_object 快 5 倍)
schema = StructType([
StructField("device_id", StringType(), True),
StructField("ts", StringType(), True),
StructField("wind_speed", DoubleType(), True),
StructField("temp_c", DoubleType(), True)
])
parsed_df = kafka_df.select(
from_json(col("value").cast("string"), schema).alias("data")
).select("data.*")
# 3. 时间戳解析(再次强调时区!)
parsed_df = parsed_df.withColumn(
"event_time",
to_timestamp(col("ts"), "yyyy-MM-dd'T'HH:mm:ssXXX")
).withColumn("event_time_utc", col("event_time"))
# 4. 设置水印(容忍 10 分钟乱序)
watermarked_df = parsed_df.withWatermark("event_time_utc", "10 minutes")
# 5. 窗口聚合(1 小时滚动窗口)
windowed_df = watermarked_df.groupBy(
window(col("event_time_utc"), "1 hour", "1 hour", "-30 minutes").alias("window"),
col("device_id")
).agg(
mean("wind_speed").alias("mean_wind_1h"),
stddev("wind_speed").alias("std_wind_1h"),
count("*").alias("sample_count")
)
# 6. 加载离线特征(用 broadcast join,避免 shuffle)
# 假设离线特征表已存为 Delta Lake
offline_features = spark.read.format("delta").load("s3a://data-lake/features/daily/")
broadcast_features = broadcast(offline_features)
# 7. 特征拼接(注意:join key 是 device_id,不是 window)
joined_df = windowed_df.join(
broadcast_features,
["device_id"],
"left"
)
# 8. 加载模型(broadcast)
with open("/tmp/model.pkl", "rb") as f:
model_bytes = f.read()
broadcast_model = spark.sparkContext.broadcast(model_bytes)
# 9. 定义预测 UDF(重点:处理 null 和边界)
@pandas_udf("struct<prediction:double, confidence_low:double, confidence_high:double>", PandasUDFType.SCALAR)
def predict_udf(
mean_wind_1h: pd.Series,
std_wind_1h: pd.Series,
sample_count: pd.Series,
lag_1h_power: pd.Series,
lag_24h_power: pd.Series,
rolling_mean_24h: pd.Series,
rolling_std_24h: pd.Series,
hour_of_day: pd.Series,
is_weekend: pd.Series,
days_since_sep: pd.Series
) -> pd.Series:
# 1. 构造特征矩阵(处理 null)
X = np.column_stack([
mean_wind_1h.fillna(0),
std_wind_1h.fillna(0),
sample_count.clip(0, 100), # 防止异常大值
lag_1h_power.fillna(method='ffill').fillna(0),
lag_24h_power.fillna(method='ffill').fillna(0),
rolling_mean_24h.fillna(method='ffill').fillna(0),
rolling_std_24h.fillna(method='ffill').fillna(0),
np.sin(2*np.pi*hour_of_day/24),
np.cos(2*np.pi*hour_of_day/24),
is_weekend.astype(int),
days_since_sep
])
# 2. 加载模型
model = pickle.loads(broadcast_model.value)
# 3. 预测(用 quantile regression 估计置信区间)
pred = model.predict(X)
# 简化版置信区间:用训练时的残差分位数
# 实际项目中,更多推荐
所有评论(0)