import jax import jax.numpy as jnp @jax.jit def lorenz63_step(state,t,params): sigma, rho, beta = params x, y, z = state dxdt = sigma * (y - x) dydt = x * (rho - z) - y dzdt = x * y - beta * z return jnp.stack([dxdt, dydt, dzdt],axis=-1) def forward_euler_step(f,y,t,params,dt): return y + dt * f(y, t, params) def rk4_step(f,y,t,params,dt): k1 = dt * f(y, t, params) k2 = dt * f(y + 0.5 * k1, t + dt/2, params) k3 = dt * f(y + 0.5 * k2, t + dt/2, params) k4 = dt * f(y + k3, t + dt, params) return y + (k1 + 2 * k2 + 2 * k3 + k4) / 6 def integrate_lorenz(params,y0,t,dt): def step(y_i,t_i): y_next = rk4_step(lorenz63_step, y_i,t_i, params, dt) return y_next, y_next _, ys = jax.lax.scan(step, y0, t[:-1]) traj = jnp.concatenate([y0[None,...], ys], axis=0) return traj