如何在已创建的Plotly Express散点图中添加趋势线?
问题解答:在Plotly Express散点图创建后添加趋势线
环境信息
- 系统:Ubuntu 20.04.3 LTS (Focal Fossa)
- Streamlit版本:1.12.0
- Python版本:3.8.10
- Plotly版本:5.10.0
问题描述
使用Plotly Express构建的Streamlit仪表盘支持用户动态选择选项重建图表,其中一个功能是添加趋势线,但目前只能在调用px.scatter()创建图表实例(fig)时通过trendline参数添加趋势线,希望能在fig = px.scatter执行完成后再为图表添加趋势线。
解决方案
Plotly Express的trendline参数是在图表初始化阶段自动完成趋势线计算与绘制的,没有直接的“事后添加”API,但可以通过手动计算趋势线拟合数据,再调用fig.add_trace()将趋势线轨迹追加到已创建的图表中,实现需求。
实现代码(修改后的函数)
首先确保安装statsmodels(Plotly Express的OLS趋势线依赖该库):
pip install statsmodels
修改后的函数逻辑:
import plotly.express as px import statsmodels.api as sm from statsmodels.formula.api import ols def control_chart_by_compound( df, x_column_name, y_column_name, trendline=False, trendlinetype="ols", trendline_scope="overall", color_column_name=None ): # 先创建基础散点图,统一逻辑避免重复 fig = px.scatter( df, x=x_column_name, y=y_column_name, labels={ "x": x_column_name, "y": y_column_name, color_column_name: "Compounds" }, color=color_column_name if color_column_name else None ) # 根据条件添加趋势线 if trendline and trendlinetype == "ols": try: if trendline_scope == "overall": # 拟合整体趋势线 model = ols(f"{y_column_name} ~ {x_column_name}", data=df).fit() df["trendline"] = model.predict(df[x_column_name]) # 添加趋势线轨迹 fig.add_trace( px.line( df, x=x_column_name, y="trendline", color_discrete_sequence=["red"] ).data[0], name="整体趋势线" ) elif trendline_scope == "trace": # 按分组添加趋势线(需要指定color_column_name) if not color_column_name: raise ValueError("trendline_scope='trace' 必须指定color_column_name参数") for group in df[color_column_name].unique(): group_df = df[df[color_column_name] == group] model = ols(f"{y_column_name} ~ {x_column_name}", data=group_df).fit() group_df["trendline"] = model.predict(group_df[x_column_name]) # 获取分组对应的散点颜色,保持视觉一致 group_color = next(trace.marker.color for trace in fig.data if trace.name == group) fig.add_trace( px.line( group_df, x=x_column_name, y="trendline", color_discrete_sequence=[group_color] ).data[0], name=f"{group} 趋势线" ) # 确保图例正常显示 fig.update_layout(showlegend=True) except Exception as e: err = f"无法为列'{x_column_name}'和'{y_column_name}'添加趋势线。(错误: {e})" print(err) return fig
关键逻辑说明
- 统一基础图表创建:先创建不带趋势线的散点图,避免原代码中重复的散点图创建逻辑。
- 趋势线计算与添加:
- 整体趋势线:用全部数据训练OLS模型,生成拟合值后添加一条线轨迹。
- 分组趋势线:按
color_column_name的分组分别训练模型,使用对应分组的散点颜色绘制趋势线,保证视觉统一。
- 轨迹追加:通过
fig.add_trace()将手动生成的趋势线轨迹添加到已有的图表实例中。
自定义扩展
- 可以修改趋势线样式,比如设置为虚线:在
px.line()中添加line_dash="dash"参数。 - 调整趋势线颜色:修改
color_discrete_sequence参数的值即可。
内容的提问来源于stack exchange,提问作者RightmireM
相关产品推荐
相关产品推荐

