大家好,我是小邓同学
今天给大家分享一个超强的算法模型,Prophet
Prophet 是 Facebook 开源的一种时间序列预测模型。
它特别适合处理具有强烈季节性波动、多重节假日效应以及包含异常值的业务数据(如电商销售额、日活跃用户数等)。
与 ARIMA、LSTM、Transformer 等模型相比,Prophet 最大的特点不是追求模型结构复杂,而是强调可解释性、鲁棒性和工程上的易用性。
Prophet 的核心思想是将时间序列分解为三个主要的业务成分:趋势项(Trend)、季节项(Seasonality)和节假日项(Holiday),再加上一个误差项。
其模型本质是一个广义可加模型
其中
Prophet 支持两种基础的趋势模型:
分段线性增长模型
若趋势增长率在某些特定时间点发生突变,其数学形式可表示为:
其中
逻辑斯谛增长模型
适用于具有饱和上限的增长场景(如用户总数增长)
其中 为承载能力, 为增长率, 为偏移参数
Prophet 利用傅里叶级数来捕捉复杂的周期性规律(例如年度季节性、周度季节性)
其中:
节假日和促销活动往往会造成时间序列出现剧烈波动。
Prophet 允许用户自定义一个包含节假日及其影响窗口的列表。
假设有 个不同的节假日,每个节假日 对应一个窗口期 。
其中
下面是一个使用 Prophet 进行股票价格预测的完整示例代码。
import warningsimport matplotlib.pyplot as pltimport numpy as npimport pandas as pdfrom prophet import Prophetimport seaborn as snsimport yfinance as yf# ==========================================# 1. 从雅虎财经获取数据 (以 AAPL 为例)# ==========================================ticker = "AAPL"# 可替换为 TSLA, NVDA 等股票代码print(f"正在从雅虎财经下载 {ticker} 历史价格数据...")# 获取过去 3 年的数据data = yf.download(ticker, start="2010-01-01", end="2026-08-01")# 处理多层索引 (yfinance 近期更新可能返回 MultiIndex 列)if isinstance(data.columns, pd.MultiIndex): data = data["Close"]else: data = data[["Close"]]# Prophet 严格要求两列数据: 'ds' (日期) 和 'y' (目标数值)df = data.reset_index()df.columns = ["ds", "y"]# 移除时区信息(Prophet 处理无时区时间序列稳定性更好)df["ds"] = pd.to_datetime(df["ds"]).dt.tz_localize(None)print(f"成功获取 {len(df)} 条交易日数据。")# ==========================================# 2. 构建并训练 Prophet 模型# ==========================================# 股票市场数据包含明显的周季节性(交易日)和年季节性model = Prophet( daily_seasonality=False, weekly_seasonality=True, yearly_seasonality=True, changepoint_prior_scale=0.05, # 变点灵活度控制,防止过拟合 seasonality_prior_scale=10.0, # 季节性先验强度 interval_width=0.95, # 设置 95% 置信区间)# 拟合模型model.fit(df)# 创建未来 180 天的预测时间序列(剔除周末以符合交易日特点)future = model.make_future_dataframe(periods=180, freq="D")future = future[future["ds"].dt.dayofweek < 5] # 保留周一至周五# 执行预测forecast = model.predict(future)# ==========================================# 3. 绘制 4 个精美画板图表# ==========================================fig = plt.figure(figsize=(16, 12), dpi=120)# ------------------------------------------# 图 1:历史收盘价与未来 Prophet 预测走势图# ------------------------------------------ax1 = plt.subplot(2, 2, 1)# 真实观测值ax1.scatter( df["ds"], df["y"], color="#1f77b4", s=12, alpha=0.6, label="Historical Close Price",)# 预测趋势线ax1.plot( forecast["ds"], forecast["yhat"], color="#d62728", lw=2, label="Forecast")# 95% 置信区间填充ax1.fill_between( forecast["ds"], forecast["yhat_lower"], forecast["yhat_upper"], color="#d62728", alpha=0.15, label="95% Confidence Interval",)ax1.set_title( f"1. {ticker} Stock Price Prediction (Prophet)", fontsize=12, fontweight="bold", pad=10,)ax1.set_ylabel("Price ($)")ax1.legend(loc="upper left", frameon=True)# ------------------------------------------# 图 2:趋势解构与变点(Changepoints)分析图# ------------------------------------------ax2 = plt.subplot(2, 2, 2)# 整体趋势线ax2.plot( forecast["ds"], forecast["trend"], color="#2ca02c", lw=2.5, label="Overall Trend",)# 标注 Prophet 自动探测的主要变点(取变化显著前10个)changepoints = model.changepointsif len(changepoints) > 0: deltas = np.abs(model.params["delta"].squeeze()) top_cps = changepoints.iloc[np.argsort(deltas)[-8:]]for cp in top_cps: ax2.axvline( x=cp, color="#ff7f0e", linestyle="--", alpha=0.6, lw=1.2, label="_nolegend_" ) ax2.axvline( x=top_cps.iloc[0], color="#ff7f0e", linestyle="--", alpha=0.6, lw=1.2, label="Significant Changepoints", )ax2.set_title("2. Underlying Trend & Detected Changepoints", fontsize=12, fontweight="bold", pad=10,)ax2.set_ylabel("Trend Magnitude")ax2.legend(loc="upper left", frameon=True)# ------------------------------------------# 图 3:每周与每年的季节性模式 (Seasonality)# ------------------------------------------ax3 = plt.subplot(2, 2, 3)# 计算一周内各天的平均影响weekly_df = forecast.groupby(forecast["ds"].dt.day_name())["weekly"].mean()days_order = ["Monday","Tuesday","Wednesday","Thursday","Friday",]weekly_df = weekly_df.reindex(days_order)# 柱状图展示周内效应colors = ["#4c72b0"if v >= 0 else"#c44e52"for v in weekly_df.values]ax3.bar( weekly_df.index, weekly_df.values, color=colors, alpha=0.85, width=0.5, edgecolor="black", lw=0.5,)ax3.axhline(0, color="gray", linewidth=0.8)ax3.set_title("3. Weekly Seasonality Pattern (Day of Week Effect)", fontsize=12, fontweight="bold", pad=10,)ax3.set_ylabel("Price Impact ($)")# ------------------------------------------# 图 4:模型训练残差(Errors)分布诊断图# ------------------------------------------ax4 = plt.subplot(2, 2, 4)# 合并真实值与拟合值,计算残差merged = pd.merge(df, forecast[["ds", "yhat"]], on="ds")residuals = merged["y"] - merged["yhat"]# 绘制残差分布核密度图sns.histplot( residuals, kde=True, ax=ax4, color="#8c564b", bins=30,stat="density", alpha=0.4,)ax4.axvline(0, color="red", linestyle="--", linewidth=1.5, label="Zero Error")ax4.set_title("4. Model Residuals Distribution (In-sample Error)", fontsize=12, fontweight="bold", pad=10,)ax4.set_xlabel("Residual (Actual - Predicted)")ax4.set_ylabel("Density")ax4.legend(loc="upper right")# 调整子图间距并展示plt.tight_layout()plt.show()
