| 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 |