如何在Matplotlib中为散点图每个点设两种颜色,同时展示真值与分类结果?
嘿,我来帮你实现把真值(ground-truth)和分类结果同时可视化的需求!你现在已经搞定了带边的t-SNE真值图,接下来只需要在现有代码基础上添加预测结果的可视化逻辑就行,这里给你几个实用的方案:
方案1:同图用不同标记区分真值与预测结果
这个方案适合快速对比同一特征点上的真值和预测类别,比如用实心圆表示真值,空心叉号表示预测结果:
from matplotlib.collections import LineCollection import matplotlib.pyplot as plt # 假设你的分类结果存储在preds变量中,先处理真值和预测值的标签与颜色 gt_lbls = data.y.cpu().numpy() - 1 pred_lbls = preds.cpu().numpy() - 1 # 确保和真值的标签偏移一致 cols_gt = ['rgbkm'[lbl] for lbl in gt_lbls] cols_pred = ['rgbkm'[lbl] for lbl in pred_lbls] # 创建边集合(和你原有代码一致) lc = LineCollection(X_embedded[out_dict['edges']], linewidth=0.05) fig = plt.figure(figsize=(10, 8)) ax = plt.gca() ax.add_collection(lc) # 绘制真值点:实心圆,半透明 ax.scatter(X_embedded[:,0], X_embedded[:,1], c=cols_gt, marker='o', alpha=0.6, label='Ground Truth') # 绘制预测结果点:空心叉号,更大尺寸,更显眼 ax.scatter(X_embedded[:,0], X_embedded[:,1], c=cols_pred, marker='x', s=50, alpha=0.8, label='Predicted') # 设置坐标轴范围(和你原有代码一致) ax.set_xlim(X_embedded[:,0].min(), X_embedded[:,0].max()) ax.set_ylim(X_embedded[:,1].min(), X_embedded[:,1].max()) # 添加图例和标题 plt.legend() plt.title('t-SNE: Ground Truth vs Predicted Labels') plt.show()
方案2:分左右子图对比(更清晰)
如果同图标记太多显得杂乱,不如直接分成两个子图,左边展示真值,右边展示预测结果:
from matplotlib.collections import LineCollection import matplotlib.pyplot as plt # 处理标签与颜色(和方案1一致) gt_lbls = data.y.cpu().numpy() - 1 pred_lbls = preds.cpu().numpy() - 1 cols_gt = ['rgbkm'[lbl] for lbl in gt_lbls] cols_pred = ['rgbkm'[lbl] for lbl in pred_lbls] lc = LineCollection(X_embedded[out_dict['edges']], linewidth=0.05) # 创建1行2列的子图布局 fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 8)) # 左子图:真值可视化 ax1.add_collection(lc) ax1.scatter(X_embedded[:,0], X_embedded[:,1], c=cols_gt) ax1.set_xlim(X_embedded[:,0].min(), X_embedded[:,0].max()) ax1.set_ylim(X_embedded[:,1].min(), X_embedded[:,1].max()) ax1.set_title('Ground Truth Labels') # 右子图:预测结果可视化 ax2.add_collection(lc) ax2.scatter(X_embedded[:,0], X_embedded[:,1], c=cols_pred) ax2.set_xlim(X_embedded[:,0].min(), X_embedded[:,0].max()) ax2.set_ylim(X_embedded[:,1].min(), X_embedded[:,1].max()) ax2.set_title('Predicted Labels') # 自动调整子图间距 plt.tight_layout() plt.show()
方案3:高亮分类错误的点
如果重点想关注哪些样本分类错了,可以把错误预测的点用特殊颜色/标记高亮:
from matplotlib.collections import LineCollection import matplotlib.pyplot as plt gt_lbls = data.y.cpu().numpy() - 1 pred_lbls = preds.cpu().numpy() - 1 cols_gt = ['rgbkm'[lbl] for lbl in gt_lbls] # 找出分类错误的样本索引 wrong_idx = gt_lbls != pred_lbls lc = LineCollection(X_embedded[out_dict['edges']], linewidth=0.05) fig = plt.figure(figsize=(10, 8)) ax = plt.gca() ax.add_collection(lc) # 绘制所有真值点 ax.scatter(X_embedded[:,0], X_embedded[:,1], c=cols_gt, alpha=0.6) # 用红色叉号高亮错误点 ax.scatter(X_embedded[wrong_idx,0], X_embedded[wrong_idx,1], c='red', s=60, marker='x', label='Misclassified') ax.set_xlim(X_embedded[:,0].min(), X_embedded[:,0].max()) ax.set_ylim(X_embedded[:,1].min(), X_embedded[:,1].max()) plt.legend() plt.title('t-SNE with Misclassified Points Highlighted') plt.show()
额外小提示
如果你的类别数超过5个,['rgbkm']的颜色就不够用了,可以改用matplotlib的内置色卡来自动生成颜色:
import matplotlib.cm as cm import numpy as np num_classes = len(np.unique(gt_lbls)) cmap = cm.get_cmap('tab10', num_classes) # tab10是常用的多类别色卡 cols_gt = [cmap(lbl) for lbl in gt_lbls] cols_pred = [cmap(lbl) for lbl in pred_lbls]
内容的提问来源于stack exchange,提问作者DsCpp
相关产品推荐
相关产品推荐

