OneOneFour's picture
latest updates except data
fce6c09
Raw
History Blame Contribute Delete
861 Bytes
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