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

使用matplotlib.pyplot绘制分组彩色散点图报错及颜色统一问题求助

Fixing the Scatter Plot Color Grouping Error in Matplotlib

Hey there, let's break down why you're seeing that ValueError and get your scatter plot showing distinct colors for your two groups!

The Root Cause

The error pops up because your y_train contains categorical group labels (like string names) instead of numerical values. Matplotlib's c parameter needs numbers to map to the colormap you chose (plt.cm.Paired)—it can't directly interpret raw category names like "sample group".

Solutions to Fix Color Grouping

1. Convert Categorical Labels to Numerical Values (Recommended)

Use LabelEncoder from scikit-learn to turn your group labels into numbers that matplotlib can work with:

from sklearn.preprocessing import LabelEncoder
import matplotlib.pyplot as plt

# Transform categorical labels to numerical codes
label_encoder = LabelEncoder()
y_train_numeric = label_encoder.fit_transform(y_train)

# Create your plot with the numeric labels
fig = plt.figure(1, figsize=(10, 6))
plt.scatter(X_train_reduced[:, 0], X_train_reduced[:, 1], 
            c=y_train_numeric, 
            cmap=plt.cm.Paired, 
            linewidths=10)

# Optional: Add a colorbar to clarify which number maps to which group
plt.colorbar(ticks=label_encoder.transform(label_encoder.classes_), 
             label='Sample Group')
plt.clim(-0.5, len(label_encoder.classes_) - 0.5) # Centers ticks on color groups
plt.show()

2. Manually Map Labels to Specific Colors

If you want full control over which color each group gets (instead of using a colormap), create a label-to-color dictionary:

import matplotlib.pyplot as plt

# Replace 'group1'/'group2' with your actual group names
color_mapping = {'group1': '#ff7f0e', 'group2': '#1f77b4'}
# Generate a color list for each data point
point_colors = [color_mapping[label] for label in y_train]

fig = plt.figure(1, figsize=(10, 6))
plt.scatter(X_train_reduced[:, 0], X_train_reduced[:, 1], 
            c=point_colors, 
            linewidths=10)
plt.show()

3. Quick Fix with Pandas factorize

If you're already using pandas, this is a fast way to convert labels to numbers:

import pandas as pd
import matplotlib.pyplot as plt

# Convert labels to numerical codes
y_train_numeric, _ = pd.factorize(y_train)

fig = plt.figure(1, figsize=(10, 6))
plt.scatter(X_train_reduced[:, 0], X_train_reduced[:, 1], 
            c=y_train_numeric, 
            cmap=plt.cm.Paired, 
            linewidths=10)
plt.show()

Quick Note on Your Code

You're calling plt.figure() twice (fig = plt.figure(1, ...) followed by plt.figure()). This creates two separate figure windows—you can remove one of these lines to keep your plot focused in a single figure.

内容的提问来源于stack exchange,提问作者user7249622

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 08:20:24