四阶龙格-库塔法求解微分方程组Python代码问题排查
问题描述
尝试用Python脚本通过四阶龙格-库塔法求解以下二元微分方程组:
dx/dt = ax - bxy dy/dt = -cy + dxy
积分区间为[15; 45],已知系数a、b、c、d及初始值x₀、y₀,但运行自行编写的代码后,绘图结果与书中示例完全不同。
示例绘图:
我的绘图:
附原始实现代码:
from typing import List, Tuple import matplotlib.pyplot as plt def get_next_v(_v_i: float, _q: List[float]) -> float: assert len(_q) == 4 _q_s = _q[0] + 2 * _q[1] + 2 * _q[2] + _q[3] return _v_i + ((1 / 6) * _q_s) def get_qi_seq(_h: float, _t_i: float, _x_i: float, _y_i: float, _a: float, _b: float, _c: float, _d: float) -> Tuple[List[float], List[float]]: # x' = ax - bxy # y` = -cy + dxy _k_1 = _h * x(_x_i, _y_i, _a, _b) _q_1 = _h * y(_x_i, _y_i, _c, _d) _k_2 = _h * x(_x_i + _h / 2, _y_i + _k_1 / 2, _a, _b) _q_2 = _h * y(_x_i + _h / 2, _y_i + _q_1 / 2, _c, _d) _k_3 = _h * x(_x_i + _h / 2, _y_i + _k_2 / 2, _a, _b) _q_3 = _h * y(_x_i + _h / 2, _y_i + _q_2 / 2, _c, _d) _k_4 = _h * x(_x_i + _h, _y_i + _k_3, _a, _b) _q_4 = _h * y(_x_i + _h, _y_i + _q_3, _c, _d) return [_k_1, _k_2, _k_3, _k_4], [_q_1, _q_2, _q_3, _q_4] def x(_x_i: float, _y_i: float, _a: float, _b: float) -> float: return (_x_i * _a) - (_b * _x_i * _y_i) def y(_x_i: float, _y_i: float, _c: float, _d: float) -> float: return (- (_y_i * _c)) + (_d * _x_i * _y_i) def get_x_y_vals(_t0: float, _T: float, _x0: float, _y0: float, _h: float, _a: float, _b: float, _c: float, _d: float) -> List[List[float]]: _t = _t0 _xi = _x0 _yi = _y0 _res = [[_t, _xi, _yi]] while _t < _T: _ki, _qi = get_qi_seq(_h, _t, _xi, _yi, _a, _b, _c, _d) _xi = get_next_v(_xi, _ki) _yi = get_next_v(_xi, _qi) _t += _h _res.append([_t, _xi, _yi]) return _res def main(): # # a = 1.75 # # b = 2.15 # # c = 1.35 # # d = 0.82 # # interval = [14, 44] # # e = 10 ** -3 # a = 1.77 b = 2.17 c = 1.38 d = 0.89 x0 = 3.39 y0 = 2.13 interval = [15, 45] e = 10 ** -3 vals = get_x_y_vals(interval[0], interval[1], x0, y0, 0.1, a, b, c, d) # print(get_h(e)) # 原代码中get_h未定义,注释掉避免报错 print(vals) t_range = [_item[0] for _item in vals] x_range = [_item[1] for _item in vals] y_range = [_item[2] for _item in vals] plt.plot(t_range, x_range) plt.plot(t_range, y_range) # plt.plot(x_range, y_range) plt.show() if __name__ == "__main__": main()
错误分析与修正
代码存在两处核心错误,直接导致计算结果偏离预期:
1. 二元变量更新顺序错误
在get_x_y_vals函数中,更新_yi时错误使用了刚更新的_xi作为初始值:
_yi = get_next_v(_xi, _qi) # 错误:第一个参数应为原_yi,而非更新后的_xi
四阶龙格-库塔法中,x和y的更新必须基于同一时刻的初始值,不能互相干扰。正确做法是先计算出下一个时刻的x和y,再同时更新:
new_x = get_next_v(_xi, _ki) new_y = get_next_v(_yi, _qi) _xi = new_x _yi = new_y
2. 龙格-库塔中间步计算错误
在get_qi_seq函数中,中间点的计算完全不符合四阶龙格-库塔法的规则:
- 错误地将步长的一半(
_h/2)直接加到x上,而非x的增量的一半(_k_1/2) - 错误地将x的增量(
_k_1)加到y上,而非y的增量(_q_1)
正确的中间点计算应该基于当前变量加上对应增量的一半,修正后的get_qi_seq函数如下:
def get_qi_seq(_h: float, _t_i: float, _x_i: float, _y_i: float, _a: float, _b: float, _c: float, _d: float) -> Tuple[List[float], List[float]]: # x' = ax - bxy = f(x,y) # y` = -cy + dxy = g(x,y) # 第一步:初始点计算 _k_1 = _h * x(_x_i, _y_i, _a, _b) _q_1 = _h * y(_x_i, _y_i, _c, _d) # 第二步:中间点(xi + k1/2, yi + q1/2) _x_mid1 = _x_i + _k_1 / 2 _y_mid1 = _y_i + _q_1 / 2 _k_2 = _h * x(_x_mid1, _y_mid1, _a, _b) _q_2 = _h * y(_x_mid1, _y_mid1, _c, _d) # 第三步:中间点(xi + k2/2, yi + q2/2) _x_mid2 = _x_i + _k_2 / 2 _y_mid2 = _y_i + _q_2 / 2 _k_3 = _h * x(_x_mid2, _y_mid2, _a, _b) _q_3 = _h * y(_x_mid2, _y_mid2, _c, _d) # 第四步:终点(xi + k3, yi + q3) _x_end = _x_i + _k_3 _y_end = _y_i + _q_3 _k_4 = _h * x(_x_end, _y_end, _a, _b) _q_4 = _h * y(_x_end, _y_end, _c, _d) return [_k_1, _k_2, _k_3, _k_4], [_q_1, _q_2, _q_3, _q_4]
修正后的完整代码
from typing import List, Tuple import matplotlib.pyplot as plt def get_next_v(_v_i: float, _q: List[float]) -> float: assert len(_q) == 4 _q_s = _q[0] + 2 * _q[1] + 2 * _q[2] + _q[3] return _v_i + ((1 / 6) * _q_s) def get_qi_seq(_h: float, _t_i: float, _x_i: float, _y_i: float, _a: float, _b: float, _c: float, _d: float) -> Tuple[List[float], List[float]]: # x' = ax - bxy = f(x,y) # y` = -cy + dxy = g(x,y) _k_1 = _h * x(_x_i, _y_i, _a, _b) _q_1 = _h * y(_x_i, _y_i, _c, _d) _x_mid1 = _x_i + _k_1 / 2 _y_mid1 = _y_i + _q_1 / 2 _k_2 = _h * x(_x_mid1, _y_mid1, _a, _b) _q_2 = _h * y(_x_mid1, _y_mid1, _c, _d) _x_mid2 = _x_i + _k_2 / 2 _y_mid2 = _y_i + _q_2 / 2 _k_3 = _h * x(_x_mid2, _y_mid2, _a, _b) _q_3 = _h * y(_x_mid2, _y_mid2, _c, _d) _x_end = _x_i + _k_3 _y_end = _y_i + _q_3 _k_4 = _h * x(_x_end, _y_end, _a, _b) _q_4 = _h * y(_x_end, _y_end, _c, _d) return [_k_1, _k_2, _k_3, _k_4], [_q_1, _q_2, _q_3, _q_4] def x(_x_i: float, _y_i: float, _a: float, _b: float) -> float: return (_x_i * _a) - (_b * _x_i * _y_i) def y(_x_i: float, _y_i: float, _c: float, _d: float) -> float: return (- (_y_i * _c)) + (_d * _x_i * _y_i) def get_x_y_vals(_t0: float, _T: float, _x0: float, _y0: float, _h: float, _a: float, _b: float, _c: float, _d: float) -> List[List[float]]: _t = _t0 _xi = _x0 _yi = _y0 _res = [[_t, _xi, _yi]] while _t < _T: _ki, _qi = get_qi_seq(_h, _t, _xi, _yi, _a, _b, _c, _d) new_x = get_next_v(_xi, _ki) new_y = get_next_v(_yi, _qi) _xi = new_x _yi = new_y _t += _h _res.append([_t, _xi, _yi]) return _res def main(): a = 1.77 b = 2.17 c = 1.38 d = 0.89 x0 = 3.39 y0 = 2.13 interval = [15, 45] e = 10 ** -3 vals = get_x_y_vals(interval[0], interval[1], x0, y0, 0.1, a, b, c, d) # print(get_h(e)) # 原函数未定义,注释避免报错 # print(vals) t_range = [_item[0] for _item in vals] x_range = [_item[1] for _item in vals] y_range = [_item[2] for _item in vals] plt.plot(t_range, x_range, label='x(t)') plt.plot(t_range, y_range, label='y(t)') plt.legend() plt.xlabel('t') plt.ylabel('Value') plt.show() if __name__ == "__main__": main()
内容的提问来源于stack exchange,提问作者Anton Grant
相关产品推荐
相关产品推荐

