如何在scipy的solve_bvp中使用随时间变化的已知参数
错误原因
solve_bvp求解过程中会动态调整计算节点,导数函数fun接收到的输入x的长度不是你初始传入的固定11个点,会根据求解需求动态变化。你直接将固定长度为11的d2_set_values赋值给d2参与运算,和x对应长度的状态变量分量形状不匹配,就会触发广播错误。
解决方案
对时间依赖参数d2做插值,根据fun当前接收的x时间点,动态从预定义的t和d2_set_values中取出对应取值,使用numpy自带的线性插值函数np.interp即可实现,该函数支持x为数组的批量计算。
如果你需要严格的分段常数取值,不需要线性插值,也可以用np.digitize做索引取值。
修改后的代码
仅需要修改fun函数中d2的赋值逻辑即可,修改后的完整函数如下:
# ODEs def fun(x, y): S1, I1, R1, S2, I2, R2, lamS1, lamI1, lamR1, lamS2, lamI2, lamR2 = y d1 = 0.5*(I1 + 0.1*I2)*(lamS1 - lamI1) # 新增插值逻辑,动态获取当前x对应的d2值,保证长度和x一致 d2 = np.interp(x, t, d2_set_values) # 分段常数取值写法: # idx = np.digitize(x, t, right=True) # d2 = d2_set_values[idx] dS1dt = -0.5*S1*(1-d1)*(I1 + 0.1*I2) dS2dt = -0.5*S2*(1-d2)*(I2 + 0.1*I1) dI1dt = 0.5*S1*(1-d1)*(I1 + 0.1*I2) - 0.2*I1 dI2dt = 0.5*S2*(1-d2)*(I2 + 0.1*I1) - 0.2*I2 dR1dt = 0.2*I1 dR2dt = 0.2*I2 dlamS1dt = 0.5*(1-d1)*S1*lamS1 dlamS2dt = 0.5*(1-d2)*S2*lamS2 dlamI1dt = 0.5*(1-d1)*I1*lamI1 dlamI2dt = 0.5*(1-d2)*I2*lamI2 dlamR1dt = lamR1 dlamR2dt = lamR2 return np.vstack((dS1dt, dI1dt, dR1dt, dS2dt, dI2dt, dR2dt, dlamS1dt, dlamI1dt, dlamR1dt, dlamS2dt, dlamI2dt, dlamR2dt))
内容的提问来源于stack exchange,提问作者darrenfwl
相关产品推荐
相关产品推荐

