如何用scipy.integrate.solve_ivp向量化求解numpy数组输入的耦合微分方程
用scipy.integrate.solve_ivp批量处理多组耦合微分方程
当你需要批量处理多组(每组包含多个耦合微分方程)的非耦合系统时,solve_ivp要求初始条件y0为一维数组,且导数函数返回的结果也必须是一维数组。你的问题核心在于没有正确处理数组的形状匹配,以下是高效的解决方案:
核心思路
将所有组的初始条件扁平化,在导数函数内部将状态变量重塑为(每组方程数, 组数)的二维结构,利用numpy向量化计算每组的导数,最后再将导数结果扁平化返回,以此避免循环,保证计算效率。
完整代码示例
import numpy as np from scipy.integrate import solve_ivp def test_fun(t, y): c = np.array([1, 2]) # 大规模c数组同样适用 # 将一维的y重塑为(3, N),N为c的长度(即组数) y_reshaped = y.reshape(3, len(c)) # 向量化计算每组的导数,每个结果形状为(N,) dy1 = 2 * t + c dy2 = 3 * t + c dy3 = 4 * t + c # 将三个导数堆叠为(3, N)后扁平化,返回一维数组 return np.vstack([dy1, dy2, dy3]).flatten() # 初始条件:每组3个变量,共2组,扁平化后为一维数组 y0 = np.array([[1, 2], [2, 2], [3, 2]]).flatten() # 求解区间[0,1] sol = solve_ivp(test_fun, [0, 1], y0) # 将最终结果重塑回(3, 2)的结构,方便查看每组的解 y_final = sol.y[:, -1].reshape(3, 2) print(y_final)
结果验证
运行代码后得到的结果与解析解完全一致:
[[3. 5. ] [4.5 5.5] [6. 6. ]]
对应每组的解析解计算:
- 第一组(c=1):
$y_1(t)=t^2 + t +1$ → $t=1$时为3
$y_2(t)=1.5t^2 + t +2$ → $t=1$时为4.5
$y_3(t)=2t^2 + t +3$ → $t=1$时为6 - 第二组(c=2):
$y_1(t)=t^2 + 2t +2$ → $t=1$时为5
$y_2(t)=1.5t^2 + 2t +2$ → $t=1$时为5.5
$y_3(t)=2t^2 + 2t +2$ → $t=1$时为6
注意事项
- 确保初始条件的扁平化顺序与导数函数内的重塑顺序一致,避免解的错位
- 所有导数计算利用numpy的向量化特性,即使
c规模很大,也能保持高效计算,无需循环处理每个元素
内容的提问来源于stack exchange,提问作者Pxxxx96
相关产品推荐
相关产品推荐

