Matplotlib折线图按条件设颜色报错:数组非有效color参数值
问题描述
我编写了如下Matplotlib绘图代码,希望依据0.05的阈值对数据进行条件格式着色:
from matplotlib.colors import to_rgba import numpy as np import matplotlib.pyplot as plt # Generate x-axis x = np.linspace(0, len(data_formula), len(data_formula)) colors = np.where(data_formula <= 0.05, "blue", "green") plt.plot(x, data_formula, c=colors) # Add labels and title plt.ylabel('Volume') plt.xlabel('time') plt.title('Energy') # Display the plot plt.show()
但出现报错:
array(['blue', 'blue', 'blue', ..., 'blue', 'blue', 'blue'], dtype='<U5') is not a valid value for color
尝试过列表等结构仍未解决,想请教问题出在哪里?color参数是否不接受某些数据结构或格式有误?
补充:data_formula定义如下:
def energy(x): return (x**2) data_formula = np.apply_along_axis(energy, axis=0, arr=data_normalized)
其数据类型为numpy.ndarray。
问题原因
Matplotlib的plt.plot()函数的c(或color)参数不支持为曲线上的每个点单独指定颜色数组——这个参数仅接受单一颜色值(比如字符串"blue"、RGB元组(0,0,1)),或者用于映射颜色条的数值数组,无法直接传递每个点对应的颜色字符串数组。
解决方案
以下是三种可行的解决方法,可根据需求选择:
方法1:拆分曲线分段绘制
通过筛选阈值对应的索引,将数据拆分为两部分分别绘制,适合需要连续曲线的场景:
import numpy as np import matplotlib.pyplot as plt # 假设data_normalized是已定义的输入数组 def energy(x): return (x**2) data_formula = np.apply_along_axis(energy, axis=0, arr=data_normalized) x = np.linspace(0, len(data_formula), len(data_formula)) # 生成阈值筛选掩码 mask = data_formula <= 0.05 # 分别绘制不同阈值区间的曲线 plt.plot(x[mask], data_formula[mask], c="blue") plt.plot(x[~mask], data_formula[~mask], c="green") # 添加标签和标题 plt.ylabel('Volume') plt.xlabel('time') plt.title('Energy') plt.show()
注:如果数据中存在非连续的阈值区间,曲线会在断点处断开,若需要连续着色曲线可参考方法3。
方法2:使用scatter绘制着色散点
如果不需要连续曲线,仅需每个点按阈值着色,可使用scatter()函数,它支持为每个点单独指定颜色:
import numpy as np import matplotlib.pyplot as plt def energy(x): return (x**2) data_formula = np.apply_along_axis(energy, axis=0, arr=data_normalized) x = np.linspace(0, len(data_formula), len(data_formula)) colors = np.where(data_formula <= 0.05, "blue", "green") plt.scatter(x, data_formula, c=colors) plt.ylabel('Volume') plt.xlabel('time') plt.title('Energy') plt.show()
方法3:使用LineCollection绘制连续着色曲线(进阶)
若需要连续曲线且每个线段片段随阈值变色,可使用LineCollection实现:
import numpy as np import matplotlib.pyplot as plt from matplotlib.collections import LineCollection def energy(x): return (x**2) data_formula = np.apply_along_axis(energy, axis=0, arr=data_normalized) x = np.linspace(0, len(data_formula), len(data_formula)) # 构造线段片段的坐标对 points = np.array([x, data_formula]).T.reshape(-1, 1, 2) segments = np.concatenate([points[:-1], points[1:]], axis=1) # 为每个线段分配颜色(基于前一个点的阈值判断) colors = np.where(data_formula[:-1] <= 0.05, "blue", "green") # 创建并添加LineCollection到坐标轴 lc = LineCollection(segments, colors=colors) fig, ax = plt.subplots() ax.add_collection(lc) ax.autoscale() # 添加标签和标题 ax.set_ylabel('Volume') ax.set_xlabel('time') ax.set_title('Energy') plt.show()
内容的提问来源于stack exchange,提问作者Infomagier
相关产品推荐
相关产品推荐

