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

在Dash中为大样本数据集实现t-SNE迭代可视化的问题

解决Dash中分步展示t-SNE计算结果的问题

要实现先展示小样本t-SNE结果、再更新全量结果的需求,核心是避免阻塞Dash主线程,并让回调能多次触发更新。可以通过后台线程+轮询组件结合dcc.Store来实现,具体步骤如下:

方案步骤

  1. 新增状态存储与轮询组件:在布局中加入dcc.Store保存计算状态和图表数据,用dcc.Interval定时检查并更新UI。
  2. 后台线程执行计算:点击启动按钮后,开启后台线程先计算小样本,再计算全量数据,过程中更新dcc.Store的状态。
  3. 轮询回调更新图表:通过dcc.Interval触发回调,根据dcc.Store中的状态返回对应图表。

完整代码示例

布局部分

import dash
from dash import html, dcc, Input, Output, State, no_update
import plotly.express as px
import pandas as pd
import threading
import json

# 假设ids是你的ID枚举类
class ids:
    START_ALGO = "start-algo"
    TSNE_GRAPH = "tsne-graph"
    DATASTORE = "datastore"
    COMPUTATION_STATE = "computation-state"
    UPDATE_INTERVAL = "update-interval"

app = dash.Dash(__name__)

app.layout = html.Div([
    html.Button("启动t-SNE计算", id=ids.START_ALGO),
    dcc.Graph(id=ids.TSNE_GRAPH),
    dcc.Store(id=ids.DATASTORE, data=pd.DataFrame().to_dict('records')),  # 存储你的数据集
    dcc.Store(id=ids.COMPUTATION_STATE, data={"status": "idle", "subset_fig": None, "full_fig": None}),
    dcc.Interval(id=ids.UPDATE_INTERVAL, interval=1000, disabled=True)  # 每秒检查一次状态
])

计算逻辑与回调

def split_data(data):
    # 实现你的数据集拆分逻辑,返回小样本和剩余数据
    df = pd.DataFrame(data)
    subset = df.sample(n=2000, random_state=42)
    rest = df.drop(subset.index)
    return subset, rest

def fast_computation_tsne(data):
    # 小样本t-SNE计算逻辑
    from openTSNE import TSNE
    tsne = TSNE(n_components=2, random_state=42, n_jobs=-1)
    return tsne.fit(data.values)

def slow_computation_tsne(data):
    # 全量t-SNE计算逻辑(可参考openTSNE的增量计算优化)
    from openTSNE import TSNE
    tsne = TSNE(n_components=2, random_state=42, n_jobs=-1)
    return tsne.fit(data.values)

def compute_tsne_background(data, store_id):
    # 后台线程执行计算
    try:
        # 1. 计算小样本并更新状态
        data_subset, data_rest = split_data(data)
        subset_res = fast_computation_tsne(data_subset)
        subset_fig = px.scatter(x=subset_res[:,0], y=subset_res[:,1], title="小样本t-SNE结果")
        
        # 更新存储状态为"已生成部分结果"
        app.server.app_context().push()
        dcc.Store(store_id).set_data({
            "status": "partial",
            "subset_fig": subset_fig.to_json(),
            "full_fig": None
        })
        
        # 2. 计算全量数据并更新状态
        full_data = pd.concat([data_subset, data_rest])
        full_res = slow_computation_tsne(full_data)
        full_fig = px.scatter(x=full_res[:,0], y=full_res[:,1], title="全量数据t-SNE结果")
        
        # 更新存储状态为"计算完成"
        dcc.Store(store_id).set_data({
            "status": "complete",
            "subset_fig": subset_fig.to_json(),
            "full_fig": full_fig.to_json()
        })
    except Exception as e:
        # 处理计算异常,更新状态
        dcc.Store(store_id).set_data({
            "status": "error",
            "error_msg": str(e),
            "subset_fig": None,
            "full_fig": None
        })

@callback(
    Output(id=ids.COMPUTATION_STATE, "data"),
    Output(id=ids.UPDATE_INTERVAL, "disabled"),
    Input(id=ids.START_ALGO, "n_clicks"),
    State(id=ids.DATASTORE, "data"),
    prevent_initial_call=True
)
def trigger_computation(n_clicks, data):
    # 启动后台计算线程
    thread = threading.Thread(
        target=compute_tsne_background,
        args=(data, ids.COMPUTATION_STATE)
    )
    thread.daemon = True
    thread.start()
    
    # 返回初始运行状态,并启用轮询
    return {"status": "running", "subset_fig": None, "full_fig": None}, False

@callback(
    Output(id=ids.TSNE_GRAPH, "figure"),
    Input(id=ids.UPDATE_INTERVAL, "n_intervals"),
    State(id=ids.COMPUTATION_STATE, "data"),
    prevent_initial_call=True
)
def update_visualization(n_intervals, computation_state):
    status = computation_state.get("status", "idle")
    
    if status == "partial":
        return json.loads(computation_state["subset_fig"])
    elif status == "complete":
        return json.loads(computation_state["full_fig"])
    elif status == "error":
        return px.scatter(title=f"计算出错: {computation_state['error_msg']}")
    else:
        raise no_update

if __name__ == "__main__":
    app.run_server(debug=True)

关键注意事项

  • 后台线程:用threading.Thread执行耗时计算,避免阻塞Dash主线程导致UI无响应。
  • 状态存储:dcc.Store用于序列化存储计算状态和图表数据,因为Plotly Figure无法直接传递,需转为JSON。
  • 上下文推送:在后台线程中更新dcc.Store时,需要调用app.server.app_context().push()确保Flask上下文有效。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 06:33:25