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

如何基于Numpy数组值切片字典,高效实现带标签颜色的散点图

解决方法

核心问题分析

  • 你之前用列表推导式循环调用plt.scatter,本质是重复绘制整个数据集10次,这是速度慢的根本原因,正确做法是单次调用scatter并为每个点指定对应颜色。
  • cdict[l_train]报错是因为numpy数组属于不可哈希类型,不能直接作为字典的键来索引。

实现步骤(纯Numpy操作,无循环)

  1. 向量化字典映射实现颜色转换:利用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()
  1. 无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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 08:25:16