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