Python中含max()的自定义函数绘制3D曲面报错如何解决
错误原因
你最初定义的fun函数仅支持接收单个标量值的a和b作为入参,但绘制3D曲面时,你传入的np.ravel(X)、np.ravel(Y)是长度为12万+的一维numpy数组。此时循环计算得到的s[i]均为和入参数组等长的数组,而非单个数值,调用Python内置的max()函数比较4个数组的大小时,解释器无法判断多个元素的数组之间的大小关系,因此抛出该错误。
解决方案
两种方案均可解决问题,按需选择即可:
方案1:修改函数适配numpy向量化操作
将返回值的内置max替换为numpy的np.max,指定按数组堆叠后的第一维取最大值,直接对整组输入并行计算,运算效率更高:
import numpy as np from mpl_toolkits.mplot3d import Axes3D import matplotlib.pyplot as plt def fun(a, b): y = np.array([2.1,2.9,3.8, 5.3]) t = np.array([0,1,2,4]) s = [] for i in range(4): s.append(abs(y[i]-a*t[i]-b)) # 沿第0维取每个坐标点对应的4个误差值的最大值 return np.max(np.array(s), axis=0) fig = plt.figure(figsize = (10, 10)) ax = fig.add_subplot(111, projection='3d') a = np.arange(-1,3,0.01) b = np.arange(-1,2,0.01) X, Y = np.meshgrid(a, b) zs = np.array(fun(np.ravel(X), np.ravel(Y))) Z= zs.reshape(X.shape) ax.plot_surface(X, Y, Z) ax.set_xlabel('a') ax.set_ylabel('b') ax.set_zlabel('f(a,b)') plt.show()
方案2:用np.vectorize包装原函数
无需修改原有函数逻辑,仅需要用numpy的vectorize方法将原函数包装为支持数组输入的版本,适合不想改动原有函数逻辑的场景:
import numpy as np from mpl_toolkits.mplot3d import Axes3D import matplotlib.pyplot as plt def fun(a, b): y=[2.1,2.9,3.8, 5.3] t=[0,1,2,4] s=[0,0,0,0] for i in range(4): s[i]=abs(y[i]-a*t[i]-b) return max(s) # 包装原函数支持数组输入 fun_vec = np.vectorize(fun) fig = plt.figure(figsize = (10, 10)) ax = fig.add_subplot(111, projection='3d') a = np.arange(-1,3,0.01) b = np.arange(-1,2,0.01) X, Y = np.meshgrid(a, b) # 调用包装后的函数 zs = np.array(fun_vec(np.ravel(X), np.ravel(Y))) Z= zs.reshape(X.shape) ax.plot_surface(X, Y, Z) ax.set_xlabel('a') ax.set_ylabel('b') ax.set_zlabel('f(a,b)') plt.show()
内容的提问来源于stack exchange,提问作者Shivaramakrishna Reddy
相关产品推荐
相关产品推荐

