如何保持热力图按actual列排序的顺序绘制prediction列热力图
解决预测类别热力图与真实类别排序对齐的问题
核心思路
先基于actual列的类别顺序确定所有轨迹段的固定排序,绘制prediction热力图时强制沿用该排序,确保两张图的轨迹段位置完全对应,实现直观对比。
修改步骤与代码实现
1. 调整函数逻辑,新增参考排序参数
修改plot_spectral函数,允许传入参考列(默认用actual)生成轨迹段排序,绘制目标列热力图时严格遵循该排序。
import seaborn as sns import matplotlib.pyplot as plt import pandas as pd def plot_spectral(df, target_col, ref_col='actual'): # 基于参考列生成轨迹段的固定排序:先按参考类别分组,组内保留原始轨迹/段顺序 sorted_idx = df.sort_values(by=[ref_col, 'traj_id', 'segment_id']).index # 按参考排序重新排列数据 aligned_df = df.loc[sorted_idx] # 构造热力图所需的结构化数据 heatmap_data = aligned_df.pivot( index=['traj_id', 'segment_id'], values=target_col ) # 绘制热力图 plt.figure(figsize=(10, 6)) sns.heatmap(heatmap_data, cmap='coolwarm', cbar=True, yticklabels=False) plt.title(f'{target_col} 热力图(与{ref_col}排序对齐)') plt.ylabel('轨迹段') plt.show()
2. 调用方式
- 绘制真实类别热力图:
plot_spectral(df, target_col='actual') - 绘制预测类别热力图(自动沿用真实类别排序):
plot_spectral(df, target_col='prediction')
关键细节说明
sorted_idx通过多列排序确保:先按参考类别分组,同一类别内按轨迹ID、段ID排序,保证排序逻辑稳定aligned_df = df.loc[sorted_idx]强制数据顺序与参考列完全匹配,两张热力图的行(轨迹段)位置一一对应- 如果需要自定义排序规则,只需修改
sort_values的by参数即可(比如调整类别内的排序优先级)
内容的提问来源于stack exchange,提问作者arilwan
相关产品推荐
相关产品推荐

