使用Matplotlib绘制密度图时遇Invalid shape错误的技术求助
问题:绘制密度图时出现形状不匹配错误
尝试使用数组x、y、z绘制密度图,但运行代码时触发形状相关错误,以下是相关代码、报错信息及预期效果:
代码
import matplotlib.pyplot as plt import matplotlib.cm as cm import numpy as np x=np.array([1., 3., 4., 5., 1., 3., 4., 5., 1., 3., 4., 5.]) y=np.array([0.1 , 0.1 , 0.1 , 0.1 , 0.01 , 0.01 , 0.01 , 0.01 , 0.001, 0.001, 0.001, 0.001]) z=np.array([-1.18976, -0.95538, -1.1647 , -1.11199, -1.02719, -0.83643, -0.94146, -0.97814, -1.27374, -1.58571, -1.65026, -1.62557]) plt.imshow(z,cmap=cm.hot,label="R0") cbar=plt.colorbar() plt.xlabel("var",size=20) plt.ylabel("\u03B2",size=20) plt.show()
报错信息
in <module> plt.imshow(z,cmap=cm.hot,label="R0") File "C:\Users\USER\anaconda3\lib\site-packages\matplotlib\pyplot.py", line 2903, in imshow __ret = gca().imshow( File "C:\Users\USER\anaconda3\lib\site-packages\matplotlib\__init__.py", line 1361, in inner return func(ax, *map(sanitize_sequence, args), **kwargs) File "C:\Users\USER\anaconda3\lib\site-packages\matplotlib\axes\_axes.py", line 5609, in imshow im.set_data(X) File "C:\Users\USER\anaconda3\lib\site-packages\matplotlib\image.py", line 709, in set_data raise TypeError("Invalid shape {} for image data" TypeError: Invalid shape (12,) for image data
预期效果
生成一个热力图,x轴标注为"var",y轴标注为"β",颜色条表示z值的大小,呈现不同x、y组合下z值的分布状态。
解决方案
错误原因
plt.imshow()要求输入的图像数据必须是二维数组(如M×N的矩阵),但当前传入的z是一维数组(12,),不符合函数的输入要求,因此触发形状错误。
修正代码
观察数据结构:x有4个唯一值(1、3、4、5),y有3个唯一值(0.1、0.01、0.001),且z的顺序是按y的每个值对应x的4个值排列的。只需将z重塑为3行4列的二维数组,再配合轴范围设置即可:
import matplotlib.pyplot as plt import matplotlib.cm as cm import numpy as np x = np.array([1., 3., 4., 5., 1., 3., 4., 5., 1., 3., 4., 5.]) y = np.array([0.1, 0.1, 0.1, 0.1, 0.01, 0.01, 0.01, 0.01, 0.001, 0.001, 0.001, 0.001]) z = np.array([-1.18976, -0.95538, -1.1647, -1.11199, -1.02719, -0.83643, -0.94146, -0.97814, -1.27374, -1.58571, -1.65026, -1.62557]) # 将一维z重塑为3行4列的二维数组,匹配y的3个层级和x的4个层级 z_2d = z.reshape(3, 4) # 获取x、y的唯一值用于轴刻度和范围 x_unique = np.unique(x) y_unique = np.unique(y) # 绘制热力图,设置轴范围并调整比例 plt.imshow(z_2d, cmap=cm.hot, extent=[x_unique.min(), x_unique.max(), y_unique.min(), y_unique.max()], aspect='auto') # 添加颜色条并设置标签 cbar = plt.colorbar(label="R0") plt.xlabel("var", size=20) plt.ylabel("β", size=20) # 设置y轴刻度为原始y的唯一值 plt.yticks(y_unique) plt.show()
说明
reshape(3,4):将一维z数组转换为符合imshow要求的二维结构,对应y的3个分组和每个分组下x的4个值。extent参数:指定x、y轴的显示范围,确保轴刻度与原始数据对应。aspect='auto':自动调整图的宽高比例,避免图形被过度拉伸或压缩。
如果更倾向于直接使用原始x、y网格绘制,也可以用pcolormesh替代imshow:
X, Y = np.meshgrid(x_unique, y_unique) plt.pcolormesh(X, Y, z_2d, cmap=cm.hot)
内容的提问来源于stack exchange,提问作者user19977266
相关产品推荐
相关产品推荐

