File size: 861 Bytes
8d4280b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 | 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 |