如何为2D散点图中的每个分组设置统一颜色
解决Matplotlib分组散点图统一颜色问题
我在使用Matplotlib绘制2D散点图时,希望为分组A、B、C…各自设置统一的颜色,但当前代码实现的是每个分组内的点颜色不同。
原代码
import matplotlib.pyplot as plt import numpy as np # Data group_labels = ['A', 'B', 'C', 'D', 'E', 'F', 'G'] data = [ [0.735721594, 0.619603837, 0.87785673, 0.482125754, 0.0894892, 0.133485767, 0.995450247], [0.666117198, 0.52923401, 0.499589112, 0.096963416, 0.308461174, 0.130418723, 0.195501054], [0.696378042, 0.437459297, 0.033071186, 0.645614608, 0.99425186, 0.097360026, 0.354376981], [0.552392974, 0.668845104, 0.079569268, 0.455465795, 0.353141333, 0.147198273, 0.249947862], [0.591065904, 0.34886412, 0.821742243, 0.008845512, 0.259947361, 0.063514992, 0.040540063], [0.016209069, 0.092671819, 0.195080351, 0.886493551, 0.745661888, 0.504613173, 0.593546542], [0.536218451, 0.466140392, 0.721903277, 0.426671591, 0.648579902, 0.823047029, 0.922809018] ] # Transpose the data to have groups on the x-axis data = np.array(data).T # Create a 2D scatter plot with unique colors for each data point color_map = plt.cm.get_cmap('tab10', len(group_labels)) # Choose a color map for i in range(len(group_labels)): colors = color_map(np.linspace(0, 1, len(data[i]))) for j in range(len(data[i])): plt.scatter(group_labels[i], data[i][j], label=f'Group {group_labels[i]}', marker='o', s=50, c=colors[j]) # Customize the plot plt.xticks(rotation=0) plt.ylabel('Y Values') plt.tight_layout() # Show the plot plt.show()
当前效果:每个分组内的点呈现不同渐变颜色,无法清晰区分分组整体。
修改后的代码
import matplotlib.pyplot as plt import numpy as np # Data group_labels = ['A', 'B', 'C', 'D', 'E', 'F', 'G'] data = [ [0.735721594, 0.619603837, 0.87785673, 0.482125754, 0.0894892, 0.133485767, 0.995450247], [0.666117198, 0.52923401, 0.499589112, 0.096963416, 0.308461174, 0.130418723, 0.195501054], [0.696378042, 0.437459297, 0.033071186, 0.645614608, 0.99425186, 0.097360026, 0.354376981], [0.552392974, 0.668845104, 0.079569268, 0.455465795, 0.353141333, 0.147198273, 0.249947862], [0.591065904, 0.34886412, 0.821742243, 0.008845512, 0.259947361, 0.063514992, 0.040540063], [0.016209069, 0.092671819, 0.195080351, 0.886493551, 0.745661888, 0.504613173, 0.593546542], [0.536218451, 0.466140392, 0.721903277, 0.426671591, 0.648579902, 0.823047029, 0.922809018] ] # Transpose the data to have groups on the x-axis data = np.array(data).T # Create a 2D scatter plot with unified color per group color_map = plt.cm.get_cmap('tab10', len(group_labels)) # Choose a color map for i in range(len(group_labels)): # 获取当前分组的单一固定颜色 group_color = color_map(i) # 一次性绘制当前分组所有点,使用统一颜色 plt.scatter(group_labels[i], data[i], label=f'Group {group_labels[i]}', marker='o', s=50, c=[group_color]*len(data[i])) # Customize the plot plt.xticks(rotation=0) plt.ylabel('Y Values') plt.legend() # 添加分组图例 plt.tight_layout() # Show the plot plt.show()
关键修改点
- 移除内层循环:不再逐个绘制点,直接对整个分组的数据批量绘制,提升代码效率
- 固定分组颜色:每个分组通过
color_map(i)获取单一颜色,而非生成渐变颜色数组 - 统一颜色传递:用
[group_color]*len(data[i])确保当前分组所有点使用同一颜色 - 修复图例:每个分组仅添加一次label,配合
plt.legend()正常显示分组图例
内容的提问来源于stack exchange,提问作者newtopy
相关产品推荐
相关产品推荐

