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

如何在已创建的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

关键逻辑说明

  1. 统一基础图表创建:先创建不带趋势线的散点图,避免原代码中重复的散点图创建逻辑。
  2. 趋势线计算与添加:
    • 整体趋势线:用全部数据训练OLS模型,生成拟合值后添加一条线轨迹。
    • 分组趋势线:按color_column_name的分组分别训练模型,使用对应分组的散点颜色绘制趋势线,保证视觉统一。
  3. 轨迹追加:通过fig.add_trace()将手动生成的趋势线轨迹添加到已有的图表实例中。

自定义扩展

  • 可以修改趋势线样式,比如设置为虚线:在px.line()中添加line_dash="dash"参数。
  • 调整趋势线颜色:修改color_discrete_sequence参数的值即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 16:30:58