如何用Matplotlib/Seaborn绘制带行列标注的模型精度矩阵
用Matplotlib和Seaborn实现精度矩阵可视化
1. 单独绘制DT与LSTM的精度矩阵
先导入所需库并准备示例数据(替换为你的真实矩阵和参数取值即可):
import matplotlib.pyplot as plt import seaborn as sns import numpy as np # 基准精度 ref_dt_accuracy = 0.86 ref_lstm_accuracy = 0.85 # 超参数取值示例 sigma_values = [0.1, 0.2, 0.3, 0.4] knots_values = [2, 3, 4, 5] # 示例精度矩阵(替换为你的真实数据) dt_accuracy_mw = np.array([ [0.84, 0.85, 0.87, 0.83], [0.85, 0.86, 0.88, 0.84], [0.83, 0.85, 0.86, 0.82], [0.82, 0.84, 0.85, 0.81] ]) lstm_accuracy_mw = np.array([ [0.83, 0.84, 0.86, 0.82], [0.84, 0.85, 0.87, 0.83], [0.82, 0.84, 0.85, 0.81], [0.81, 0.83, 0.84, 0.80] ])
接着绘制两个子图,分别展示不同超参数组合下的DT和LSTM精度:
# 创建画布和子图 fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 5)) # 绘制DT精度矩阵 sns.heatmap(dt_accuracy_mw, ax=ax1, annot=True, fmt=".2f", cmap="Blues", xticklabels=knots_values, yticklabels=sigma_values, cbar=True) ax1.set_title("DT Model Accuracy (Magnitude Warping)") ax1.set_xlabel("Knots") ax1.set_ylabel("Sigma") # 绘制LSTM精度矩阵 sns.heatmap(lstm_accuracy_mw, ax=ax2, annot=True, fmt=".2f", cmap="Greens", xticklabels=knots_values, yticklabels=sigma_values, cbar=True) ax2.set_title("LSTM Model Accuracy (Magnitude Warping)") ax2.set_xlabel("Knots") ax2.set_ylabel("Sigma") # 调整布局避免重叠 plt.tight_layout() plt.show()
生成的热力图会以颜色深浅反映精度高低,每个单元格显示具体数值,行对应sigma取值,列对应knots取值。
2. 绘制包含精度与基准差值的组合矩阵
先构造每个单元格的文本内容,再结合热力图展示双重信息:
# 构造组合注释文本 annotations = [] for dt_row, lstm_row in zip(dt_accuracy_mw, lstm_accuracy_mw): row_annot = [] for dt_acc, lstm_acc in zip(dt_row, lstm_row): dt_diff = ref_dt_accuracy - dt_acc lstm_diff = ref_lstm_accuracy - lstm_acc # 按要求格式化文本 text = f"{dt_acc:.2f} ({dt_diff:.2f})\n/{lstm_acc:.2f} ({lstm_diff:.2f})" row_annot.append(text) annotations.append(row_annot) # 绘制组合矩阵 plt.figure(figsize=(10, 8)) # 这里用DT精度作为热力图颜色依据,也可替换为lstm_accuracy_mw sns.heatmap(dt_accuracy_mw, annot=np.array(annotations), fmt="", cmap="Purples", xticklabels=knots_values, yticklabels=sigma_values, cbar=True) plt.title("DT & LSTM Accuracy with Baseline Difference") plt.xlabel("Knots") plt.ylabel("Sigma") plt.tight_layout() plt.show()
该热力图的每个单元格会同时展示两个模型的精度及与基准的差值,背景颜色直观反映对应模型的精度水平,方便对比超参数对两类模型的影响。
内容的提问来源于stack exchange,提问作者Unistack
相关产品推荐
相关产品推荐

