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

使用Numba加速RK4时出现TypingError的原因及解决方法

错误原因分析
  • 核心问题是测试函数fff返回Python列表,而RK4求解器中hdt * k1这类操作要求操作数为numpy数值类型(此处hdt是float64),Numba的nopython模式不支持Python列表与numpy标量的直接乘法,触发类型不匹配错误。
  • 额外问题:fff的输出维度和输入X0(2维数组)不匹配,即便类型修复,后续也会出现维度错误。
解决方案

修改测试函数fff,使其返回与输入维度一致的numpy数组,同时保持浮点类型:

@numba.jit(nopython=True)
def fff(t, X):
    # 返回和X同维度的浮点数组,逻辑可按需修改
    res = np.zeros_like(X, dtype=np.float64)
    res[0] = 1 * 3  # 对应原逻辑的t*X
    res[1] = 0  # 补充第二个维度的值,匹配输入X的维度
    return res

确保RK4调用时输入维度和函数输出维度一致:主运行代码中X0是2维数组,fff必须返回2维数组。

验证修改后的完整代码

swing.py修改后代码

import numpy as np
import numba

@numba.jit(nopython=True)
def RK4(func, t_end, X0, dt):
    t = np.arange(0,t_end, dt, dtype=np.float64)
    X  = np.zeros((t.shape[0], X0.shape[0]))
    X[0] = X0
    hdt = dt*.5
    for i in range(t.shape[0]-1):
        t1 = t[i]
        x1 = X[i]
        k1 = func(t[i], X[i])
        
        t2 = t[i] + hdt
        x2 = X[i] + hdt * k1
        k2 = func(t2, x2)
        
        t3 = t[i] + hdt
        x3 = X[i] + hdt * k2
        k3 = func(t3, x3)
        
        t4 = t[i] + dt
        x4 = X[i] + dt * k3
        k4 = func(t4, x4)
        X[i+1] = X[i] + dt / 6. * (k1 + 2. * k2 + 2. * k3 + k4)
    return X

# 修复后的测试函数
@numba.jit(nopython=True)
def fff(t, X):
    res = np.zeros_like(X, dtype=np.float64)
    res[0] = 1 * 3
    res[1] = 0
    return res

主运行代码

import numpy as np
import numba
import swing

swing.RK4(swing.fff, 10, np.array([0,1]), 0.1)

内容的提问来源于stack exchange,提问作者Young Jin Kim

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 14:22:01