如何基于Numpy数组值切片字典,高效实现带标签颜色的散点图
解决方法
核心问题分析
- 你之前用列表推导式循环调用
plt.scatter,本质是重复绘制整个数据集10次,这是速度慢的根本原因,正确做法是单次调用scatter并为每个点指定对应颜色。 cdict[l_train]报错是因为numpy数组属于不可哈希类型,不能直接作为字典的键来索引。
实现步骤(纯Numpy操作,无循环)
- 向量化字典映射实现颜色转换:利用
np.vectorize将字典的映射逻辑转为向量化操作,直接作用于标签数组:
import numpy as np import matplotlib.pyplot as plt # 定义你的数据 m_hat = np.array([ [17.574, 17.8316], [22.449, 23.0995], [13.4923, 11.8801], [8.34949, 8.0102], [16.676, 17.2908], [24.8699, 25.2985], [13.7985, 12.8801], [13.4541, 13.9107], [14.7577, 14.9133], [47.0102, 48.4668] ]) cdict = {1: 'red', 3: 'green', 5: 'blue', 7: 'yellow'} l_train = np.array([1,1,1,1,1,1,1,1,1,1]) # 生成每个标签对应的颜色数组 color_mapper = np.vectorize(cdict.get) point_colors = color_mapper(l_train) # 单次绘制散点图,指定所有点的颜色 plt.scatter(m_hat[:, 0], m_hat[:, 1], c=point_colors) plt.show()
- 无
np.vectorize替代方案:通过Numpy数组索引实现标签到颜色的映射:
# 将字典的键和值转为Numpy数组 keys = np.array(list(cdict.keys())) values = np.array(list(cdict.values())) # 匹配每个标签在键数组中的位置,索引对应颜色 indices = np.searchsorted(keys, l_train) point_colors = values[indices] # 绘制散点图 plt.scatter(m_hat[:, 0], m_hat[:, 1], c=point_colors) plt.show()
关键说明
- 两种方法都依赖Numpy的向量化操作,完全避免了显式循环。
- 仅需调用一次
plt.scatter即可完成所有点的上色绘制,效率远高于循环调用。
内容的提问来源于stack exchange,提问作者Karantai
相关产品推荐
相关产品推荐

