Python绘制3D图时x/y维度不匹配ValueError报错解决
错误原因
报错由两个问题共同导致:
- 传入3D绘图接口的数组长度不匹配:原代码中
K、E在内层K1循环中追加值,总长度为784,但N在外层N1循环中追加值,长度仅为28,三者维度不一致无法绘图。 - 绘图方法选择错误:
ax.plot()用于绘制3D空间折线,无法展示双参数网格对应的误差分布,无法满足3D误差可视化的需求。
另外原代码的参数遍历范围有遗漏:Pythonrange()是左闭右开区间,原写法range(3,31,1)最大只能取到30、range(20,300,10)最大只能取到290,不符合设定的参数范围要求。
修复方案
按以下步骤调整即可正常出图:
- 修正参数生成逻辑,保证K1覆盖20300、N1覆盖331的全范围
- 调整值追加逻辑,在内层循环同步追加K1、N1、误差值,保证三个一维数组长度完全一致
- 用
np.meshgrid将一维参数序列转为二维网格矩阵,将误差数组reshape为对应维度 - 替换
ax.plot()为ax.plot_surface()绘制误差曲面,可直观展示两个参数变化时误差的连续变化趋势
修复后的可运行代码如下:
import numpy as np import scipy.stats as si from scipy import integrate import matplotlib.pyplot as plt from mpl_toolkits.mplot3d import Axes3D r = 0.03 S0 = 100 T = 1 u1 = 1/5 mu = r strike = 100 sigma = 0.3 # 生成符合范围要求的参数序列 K1_list = np.arange(20, 310, 10) N1_list = np.arange(3, 32, 1) # 预计算基准期权价值 d1 = (np.log(S0 / strike) + (r + 0.5 * sigma ** 2) * T) / (sigma * np.sqrt(T)) d2 = (np.log(S0 / strike) + (r - 0.5 * sigma ** 2) * T) / (sigma * np.sqrt(T)) call = S0 * si.norm.cdf(d1, 0.0, 1.0) - strike * np.exp(-r * T) * si.norm.cdf(d2, 0.0, 1.0) K = [] N = [] E = [] for N1 in N1_list: for K1 in K1_list: def f(x): d = (np.log(x/strike) + (r + 0.5 * sigma ** 2) * (T-u1)) / (sigma * np.sqrt(T-u1)) W1 = si.norm.pdf(d) / (x*sigma*np.sqrt(T-u1)) d11 = (np.log(S0/x) + (r + 0.5 * sigma ** 2) * u1) / (sigma * np.sqrt(u1)) d22 = (np.log(S0/x) + (r - 0.5 * sigma ** 2) * u1) / (sigma * np.sqrt(u1)) call1 = S0 * si.norm.cdf(d11, 0.0, 1.0) - x * np.exp(-r * u1) * si.norm.cdf(d22, 0.0, 1.0) return W1 * call1 hedge = integrate.fixed_quad(f, 0, K1, n=N1) error = np.log(np.abs(hedge[0] - call)) # 三个值同步追加,保证维度匹配 K.append(K1) N.append(N1) E.append(error) # 转换为numpy数组并生成绘图网格 K = np.array(K) N = np.array(N) E = np.array(E) K_grid, N_grid = np.meshgrid(K1_list, N1_list) E_grid = E.reshape(len(N1_list), len(K1_list)) # 绘制3D误差曲面 fig = plt.figure(figsize=(10,10)) ax = fig.add_subplot(111, projection='3d') ax.set_title("Error in option value") ax.set_xlabel("Value of K1") ax.set_ylabel("Number of quadrature points") ax.set_zlabel("Log error") surf = ax.plot_surface(K_grid, N_grid, E_grid, cmap='viridis', edgecolor='none') fig.colorbar(surf, shrink=0.5, aspect=5, label='Log error') plt.savefig("some.png", dpi=200) plt.show()
如果你不需要连续曲面,也可以将曲面绘制语句替换为
ax.scatter(K, N, E)绘制三维散点图,同样可以展示所有参数组合对应的误差值。
内容的提问来源于stack exchange,提问作者Purba Banerjee
相关产品推荐
相关产品推荐

