You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

四阶龙格-库塔法求解微分方程组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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.08 12:15:39