如何添加带额外信息的彩色编码线条并优化绘图运行效率
问题描述
我计划绘制列车**悬挂行程(Suspension Travel)与距离(Distance)**的关系图,现有数据集包含这两项数据,以及列车处于直线轨道还是弯道的标识信息。我希望复刻下图样式,并额外添加桥梁位置、隧道等信息,但当前实现方案运行耗时近4分钟。
目标参考图
当前实现生成的图
当前使用的代码如下:
# 绘制想要添加的轨道类型标识线 def transition_line(xmin, xmax): for i in range((df_irv['Distance'] - xmin).abs().idxmin(), (df_irv['Distance'] - xmax).abs().idxmin()): plt.plot([df_irv['Distance'][i], df_irv['Distance'][i+1]], [max(df_irv['SuspTravel']) + 0.5]*2, color='red' if df_irv['Element'][i] == 'CURVA' else 'blue', linewidth=10, alpha=0.5) # 绘制带可调x轴范围的图表函数 def plot_graph(xmin, xmax, sensors): plt.figure(figsize=(10, 5)) plt.plot(df_irv['Distance'], df_irv[sensors], label='悬挂行程传感器') plt.xlim(xmin, xmax) plt.xlabel('距离 (Km)') plt.ylabel('悬挂行程传感器') plt.title('悬挂行程传感器 vs 距离') plt.legend() plt.grid(True) transition_line(xmin, xmax) plt.show() # 创建x轴范围调节滑块 xmin_slider = IntSlider(value=0, min=0, max=df_irv['Distance'].max(), step=1, description='X最小值') xmax_slider = IntSlider(value=20, min=0, max=df_irv['Distance'].max(), step=1, description='X最大值') # 交互式绘图 interact(plot_graph, xmin=xmin_slider, xmax=xmax_slider, sensors = ['SuspTravel', 'Roll', 'Bounce'])
优化方案
1. 核心性能优化:替换逐段绘制逻辑
当前transition_line用循环逐段绘制是性能瓶颈,改用broken_barh批量绘制色块,大幅减少绘图调用次数:
def transition_line(xmin, xmax, ax): max_susp = df_irv['SuspTravel'].max() y_val = max_susp + 0.5 y_height = 0.2 # 筛选x轴范围内的数据并合并连续同类型轨道区间 df_sub = df_irv[(df_irv['Distance'] >= xmin) & (df_irv['Distance'] <= xmax)].copy() if len(df_sub) == 0: return df_sub['group'] = (df_sub['Element'] != df_sub['Element'].shift()).cumsum() grouped = df_sub.groupby('group') for _, group in grouped: start = group['Distance'].iloc[0] end = group['Distance'].iloc[-1] color = 'red' if group['Element'].iloc[0] == 'CURVA' else 'blue' ax.broken_barh([(start, end - start)], (y_val - y_height/2, y_height), facecolors=color, alpha=0.5)
2. 优化交互式绘图逻辑
避免每次交互都重新创建图表,只更新内容:
import ipywidgets as widgets import numpy as np # 提前创建固定图表框架 fig, ax = plt.subplots(figsize=(10, 5)) line, = ax.plot(df_irv['Distance'], df_irv['SuspTravel'], label='悬挂行程传感器') ax.set_xlabel('距离 (Km)') ax.set_ylabel('悬挂行程传感器') ax.set_title('悬挂行程传感器 vs 距离') ax.grid(True) max_susp = df_irv['SuspTravel'].max() def update_plot(xmin, xmax, sensors): # 更新主曲线数据 line.set_data(df_irv['Distance'], df_irv[sensors]) # 调整轴范围 ax.set_xlim(xmin, xmax) ax.set_ylim(df_irv[sensors].min() - 0.5, max_susp + 1.5) # 清除旧的轨道标识 for patch in ax.patches: patch.remove() # 绘制新的轨道标识 transition_line(xmin, xmax, ax) # 添加桥梁/隧道标识(需数据集对应字段) plot_infrastructure(xmin, xmax, ax) # 更新图例 handles, labels = ax.get_legend_handles_labels() by_label = dict(zip(labels, handles)) ax.legend(by_label.values(), by_label.keys()) fig.canvas.draw_idle() # 创建交互控件 xmin_slider = widgets.IntSlider(value=0, min=0, max=int(df_irv['Distance'].max()), step=1, description='X最小值') xmax_slider = widgets.IntSlider(value=20, min=0, max=int(df_irv['Distance'].max()), step=1, description='X最大值') sensor_dropdown = widgets.Dropdown(options=['SuspTravel', 'Roll', 'Bounce'], value='SuspTravel', description='传感器') # 绑定控件与更新函数 ui = widgets.HBox([xmin_slider, xmax_slider, sensor_dropdown]) out = widgets.interactive_output(update_plot, {'xmin': xmin_slider, 'xmax': xmax_slider, 'sensors': sensor_dropdown}) display(ui, out)
3. 添加桥梁、隧道标识
参照轨道标识逻辑,用不同颜色色块展示基础设施:
def plot_infrastructure(xmin, xmax, ax): max_susp = df_irv['SuspTravel'].max() y_infra = max_susp + 1.2 infra_height = 0.2 # 绘制桥梁(假设数据集有Bridge字段标记) df_bridge = df_irv[(df_irv['Distance'] >= xmin) & (df_irv['Distance'] <= xmax) & (df_irv['Bridge'] == True)] if len(df_bridge) > 0: df_bridge['group'] = (df_bridge['Bridge'] != df_bridge['Bridge'].shift()).cumsum() for _, group in df_bridge.groupby('group'): start = group['Distance'].iloc[0] end = group['Distance'].iloc[-1] ax.broken_barh([(start, end - start)], (y_infra - infra_height/2, infra_height), facecolors='green', alpha=0.5, label='桥梁') # 绘制隧道(假设数据集有Tunnel字段标记) df_tunnel = df_irv[(df_irv['Distance'] >= xmin) & (df_irv['Distance'] <= xmax) & (df_irv['Tunnel'] == True)] if len(df_tunnel) > 0: df_tunnel['group'] = (df_tunnel['Tunnel'] != df_tunnel['Tunnel'].shift()).cumsum() for _, group in df_tunnel.groupby('group'): start = group['Distance'].iloc[0] end = group['Distance'].iloc[-1] ax.broken_barh([(start, end - start)], (y_infra - infra_height/2, infra_height), facecolors='gray', alpha=0.5, label='隧道')
内容的提问来源于stack exchange,提问作者Tiago Amorim
相关产品推荐
相关产品推荐

