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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 10:17:27