如何在满足条件时终止scipy.integrate.odeint的ODE积分?
问题
我正在使用scipy.integrate.odeint模拟双摆的翻转时间。当前代码通过odeint在固定时间区间求解系统,之后再检查是否发生翻转。我希望改为在每一步积分后检查翻转条件(angle2的变化量≥2π),一旦满足立即停止积分,无需求解剩余时间区间。现有代码如下:
from matplotlib.colors import LogNorm, Normalize from scipy.integrate import odeint import matplotlib.pyplot as plt from tqdm import tqdm import seaborn as sns import numpy as np def ode(y, t, length_1, length_2, mass_1, mass_2, gravity): angle_1, angle_1_d, angle_2, angle_2_d = y angle_1_dd = (-gravity * (2 * mass_1 + mass_2) * np.sin(angle_1) - mass_2 * gravity * np.sin(angle_1 - 2 * angle_2) - 2 * np.sin(angle_1 - angle_2) * mass_2 * (angle_2_d ** 2 * length_2 + angle_1_d ** 2 * length_1 * np.cos(angle_1 - angle_2))) / (length_1 * (2 * mass_1 + mass_2 - mass_2 * np.cos(2 * angle_1 - 2 * angle_2))) angle_2_dd = (2 * np.sin(angle_1 - angle_2) * (angle_1_d ** 2 * length_1 * (mass_1 + mass_2) + gravity * (mass_1 + mass_2) * np.cos(angle_1) + angle_2_d ** 2 * length_2 * mass_2 * np.cos(angle_1 - angle_2))) / (length_2 * (2 * mass_1 + mass_2 - mass_2 * np.cos(2 * angle_1 - 2 * angle_2))) return [angle_1_d, angle_1_dd, angle_2_d, angle_2_dd] def double_pendulum(length_1, length_2, mass_1, mass_2, angle_1_init, angle_2_init, angle_1_d_init, angle_2_d_init, gravity, dt, num_steps): time_span = np.linspace(0, dt*num_steps, num_steps) y0 = [np.deg2rad(angle_1_init), np.deg2rad(angle_1_d_init), np.deg2rad(angle_2_init), np.deg2rad(angle_2_d_init)] sol = odeint(ode, y0, time_span, args=(length_1, length_2, mass_1, mass_2, gravity)) return sol def flip(length_1, length_2, mass_1, mass_2, angle_1_init, angle_2_init, angle_1_d_init, angle_2_d_init, gravity, dt, num_steps): solution = double_pendulum(length_1, length_2, mass_1, mass_2, angle_1_init, angle_2_init, angle_1_d_init, angle_2_d_init, gravity, dt, num_steps) #angle_1 = solution[:, 0] angle_2 = solution[:, 2] for index, angle2 in enumerate(angle_2): if abs(angle2 - angle_2_init) >= 2*np.pi: return index*dt return dt*num_steps angle_1_range = np.arange(-172, 172, 1) angle_2_range = np.arange(-172, 172, 1) fliptime_matrix = np.zeros((len(angle_1_range), len(angle_2_range))) for i, angle_1 in tqdm(enumerate(angle_1_range), desc='angle_1'): for j, angle_2 in tqdm(enumerate(angle_2_range), desc='angle_2', leave=False): fliptime = flip(1, 1, 1, 1, angle_1, angle_2, 0, 0, 9.81, 0.01, 10000) fliptime_matrix[i, j] = fliptime sns.heatmap(fliptime_matrix, square=True, cbar_kws={'label': 'Divergence'}, norm=LogNorm()) plt.xlabel('Angle 2 (degrees)') plt.ylabel('Angle 1 (degrees)') plt.title('Fliptime Heatmap') plt.gca().invert_yaxis() plt.show()
解决方案
odeint本身不支持中途停止积分的功能,想要实现每步检查并提前终止,推荐改用scipy.integrate.solve_ivp——它支持事件触发终止的特性,正好匹配你的需求。具体实现如下:
from matplotlib.colors import LogNorm, Normalize from scipy.integrate import solve_ivp import matplotlib.pyplot as plt from tqdm import tqdm import seaborn as sns import numpy as np def ode(t, y, length_1, length_2, mass_1, mass_2, gravity): # solve_ivp要求ODE函数第一个参数是时间t,第二个是状态y,和odeint参数顺序相反 angle_1, angle_1_d, angle_2, angle_2_d = y angle_1_dd = (-gravity * (2 * mass_1 + mass_2) * np.sin(angle_1) - mass_2 * gravity * np.sin(angle_1 - 2 * angle_2) - 2 * np.sin(angle_1 - angle_2) * mass_2 * (angle_2_d ** 2 * length_2 + angle_1_d ** 2 * length_1 * np.cos(angle_1 - angle_2))) / (length_1 * (2 * mass_1 + mass_2 - mass_2 * np.cos(2 * angle_1 - 2 * angle_2))) angle_2_dd = (2 * np.sin(angle_1 - angle_2) * (angle_1_d ** 2 * length_1 * (mass_1 + mass_2) + gravity * (mass_1 + mass_2) * np.cos(angle_1) + angle_2_d ** 2 * length_2 * mass_2 * np.cos(angle_1 - angle_2))) / (length_2 * (2 * mass_1 + mass_2 - mass_2 * np.cos(2 * angle_1 - 2 * angle_2))) return [angle_1_d, angle_1_dd, angle_2_d, angle_2_dd] def flip_event(t, y, length_1, length_2, mass_1, mass_2, gravity, angle_2_init_rad): # 事件函数:当angle2与初始值的差值绝对值≥2π时返回0,触发终止 current_angle2 = y[2] return abs(current_angle2 - angle_2_init_rad) - 2*np.pi # 设置事件属性:触发时终止积分,direction=0表示正负方向翻转都触发 flip_event.terminal = True flip_event.direction = 0 def flip(length_1, length_2, mass_1, mass_2, angle_1_init, angle_2_init, angle_1_d_init, angle_2_d_init, gravity, max_time): angle_1_init_rad = np.deg2rad(angle_1_init) angle_2_init_rad = np.deg2rad(angle_2_init) angle_1_d_init_rad = np.deg2rad(angle_1_d_init) angle_2_d_init_rad = np.deg2rad(angle_2_d_init) y0 = [angle_1_init_rad, angle_1_d_init_rad, angle_2_init_rad, angle_2_d_init_rad] # 调用solve_ivp并传入事件函数 sol = solve_ivp(ode, t_span=[0, max_time], y0=y0, args=(length_1, length_2, mass_1, mass_2, gravity), events=flip_event, event_args=(angle_2_init_rad,)) # 返回触发事件的时间,未触发则返回最大时间 if sol.t_events[0].size > 0: return sol.t_events[0][0] else: return max_time angle_1_range = np.arange(-172, 172, 1) angle_2_range = np.arange(-172, 172, 1) fliptime_matrix = np.zeros((len(angle_1_range), len(angle_2_range))) max_time = 0.01 * 10000 # 对应原代码的总时间上限 for i, angle_1 in tqdm(enumerate(angle_1_range), desc='angle_1'): for j, angle_2 in tqdm(enumerate(angle_2_range), desc='angle_2', leave=False): fliptime = flip(1, 1, 1, 1, angle_1, angle_2, 0, 0, 9.81, max_time) fliptime_matrix[i, j] = fliptime sns.heatmap(fliptime_matrix, square=True, cbar_kws={'label': 'Divergence'}, norm=LogNorm()) plt.xlabel('Angle 2 (degrees)') plt.ylabel('Angle 1 (degrees)') plt.title('Fliptime Heatmap') plt.gca().invert_yaxis() plt.show()
关键修改说明
- ODE函数参数调整:
solve_ivp要求时间t作为第一个参数,状态y作为第二个,因此调整了ode函数的参数顺序 - 事件函数逻辑:
flip_event专门检测翻转条件,返回值为0时触发终止;terminal=True确保触发时立即停止积分,direction=0允许顺时针/逆时针翻转都触发 - 求解效率优化:无需预先生成所有时间步,求解器会自动处理步长,满足条件后立刻停止,大幅减少不必要的计算
- 参数传递:通过
event_args将初始angle2的弧度值传递给事件函数,保证条件判断的准确性
内容的提问来源于stack exchange,提问作者user1254621
相关产品推荐
相关产品推荐

