Matplotlib中for循环绘制多函数报错:k≥0时x与y维度不匹配
问题:k≥0时循环绘图出现维度不匹配错误
使用for循环为每个参数k绘制曲线,所有负k值的曲线均可正常生成,但循环到k=0或更大值时,lambdify函数抛出x与y维度不匹配的错误。单独绘制任意k值曲线均正常,但循环迭代12次后崩溃。
原代码
import sympy as sym import numpy as np import matplotlib.pyplot as plt eta = np.logspace(-1,2,21) #defines eta values, 21 decades from 0.1 to 100 relrho = np.logspace(-2,2,25) #defines values of rho2/rho1, 25 values from 0.01 to 100 k = (relrho-1)/(relrho+1) #defines the reflection coefficient #parameter of type curve is k #rhoa/rho1 is the y-axis #eta is the x-axis #R is assigned as the ratio of rho_a to rho_1 #x is assigned to eta #y is assigned to k x = sym.symbols('x', real = True) y = sym.symbols('y') for y in k: #for-loop assumes k value before while-loop is run, then plots the curve, then new k value is assumed n=1; R=1; while n<=500: Rnew = 2*x**3*y**n/(((2*n)**2+x**2)**(3/2)) R = R + Rnew n = n + 1 R = sym.lambdify(x,R) plt.loglog(eta, R(eta)) plt.show()
报错信息
runfile('C:/Users/aslak/OneDrive/Desktop/Typecurves.py', wdir='C:/Users/aslak/OneDrive/Desktop') Traceback (most recent call last): File "C:\Users\aslak\OneDrive\Desktop\Typecurves.py", line 44, in <module> plt.loglog(eta, R(eta)) File "C:\Users\aslak\anaconda3\lib\site-packages\matplotlib\pyplot.py", line 2750, in loglog return gca().loglog(*args, **kwargs) File "C:\Users\aslak\anaconda3\lib\site-packages\matplotlib\axes\_axes.py", line 1868, in loglog return self.plot( File "C:\Users\aslak\anaconda3\lib\site-packages\matplotlib\axes\_axes.py", line 1743, in plot lines = [*self._get_lines(*args, data=data, **kwargs)] File "C:\Users\aslak\anaconda3\lib\site-packages\matplotlib\axes\_base.py", line 273, in __call__ yield from self._plot_args(this, kwargs) File "C:\Users\aslak\anaconda3\lib\site-packages\matplotlib\axes\_base.py", line 399, in _plot_args raise ValueError(f"x and y must have same first dimension, but " ValueError: x and y must have same first dimension, but have shapes (21,) and (1,) runfile('C:/Users/aslak/OneDrive/Desktop/Typecurves.py', wdir='C:/Users/aslak/OneDrive/Desktop') Traceback (most recent call last): File "C:\Users\aslak\OneDrive\Desktop\Typecurves.py", line 34, in <module> plt.loglog(eta, R(eta)) File "C:\Users\aslak\anaconda3\lib\site-packages\matplotlib\pyplot.py", line 2750, in loglog return gca().loglog(*args, **kwargs) File "C:\Users\aslak\anaconda3\lib\site-packages\matplotlib\axes\_axes.py", line 1868, in loglog return self.plot( File "C:\Users\aslak\anaconda3\lib\site-packages\matplotlib\axes\_axes.py", line 1743, in plot lines = [*self._get_lines(*args, data=data, **kwargs)] File "C:\Users\aslak\anaconda3\lib\site-packages\matplotlib\axes\_base.py", line 273, in __call__ yield from self._plot_args(this, kwargs) File "C:\Users\aslak\anaconda3\lib\site-packages\matplotlib\axes\_base.py", line 399, in _plot_args raise ValueError(f"x and y must have same first dimension, but " ValueError: x and y must have same first dimension, but have shapes (21,) and (1,)
解决方案及修改后代码
问题根源
- 变量名冲突:循环中用
y作为迭代变量,覆盖了之前定义的sympy符号y = sym.symbols('y'),导致后续符号计算逻辑混乱。 - 浮点数幂运算问题:代码中
(3/2)是Python原生浮点数,sympy处理时会丢失符号类型,导致lambdify生成的函数无法正确对numpy数组做向量化运算,当k≥0时这个问题暴露出来。 - while循环可靠性低:手动while循环累加容易出现隐形错误,改用sympy原生求和函数更稳定。
修改后代码(推荐用求和函数)
import sympy as sym import numpy as np import matplotlib.pyplot as plt eta = np.logspace(-1, 2, 21) relrho = np.logspace(-2, 2, 25) k = (relrho - 1)/(relrho + 1) x = sym.symbols('x', real=True) for current_k in k: # 定义求和项,用sympy符号保持运算精度 n_sym = sym.Symbol('n') term = 2 * x**3 * current_k**n_sym / (( (2*n_sym)**2 + x**2 )**sym.Rational(3, 2)) # 求和n从1到500,加上初始值1得到完整表达式 R_expr = 1 + sym.summation(term, (n_sym, 1, 500)) # 指定numpy后端,确保向量化运算支持 R = sym.lambdify(x, R_expr, 'numpy') plt.loglog(eta, R(eta)) plt.show()
保留while循环的修改版本
import sympy as sym import numpy as np import matplotlib.pyplot as plt eta = np.logspace(-1,2,21) relrho = np.logspace(-2,2,25) k = (relrho-1)/(relrho+1) x = sym.symbols('x', real = True) for current_k in k: # 更换循环变量名,避免覆盖sympy符号 n=1; R=1; while n<=500: # 用sym.S(3)/2替代3/2,保持符号类型 Rnew = 2*x**3*current_k**n/(((2*n)**2+x**2)**(sym.S(3)/2)) R = R + Rnew n = n + 1 # 指定numpy后端保证向量化 R = sym.lambdify(x, R, 'numpy') plt.loglog(eta, R(eta)) plt.show()
内容的提问来源于stack exchange,提问作者Aslak Holm Brunvand
相关产品推荐
相关产品推荐

