@@ -52,94 +52,92 @@ def body(state):
5252
5353
5454def _theta_dy (f , j , h , th , y0 , y1 , f_y0 , pars , tol = 1e-4 ):
55- J = np .eye (y1 .size ) - h * th * j (y1 , * pars )
56- F = y0 + h * (th * f (y1 , * pars ) + f_y0 ) - y1
55+ J = np .eye (y1 .size ) - h * th * j (y1 , pars )
56+ F = y0 + h * (th * f (y1 , pars ) + f_y0 ) - y1
5757
5858 dy , _ = jacobi (J , F , w = 2.0 / 3.0 , tol = tol , max_iters = 10 )
5959 return dy
6060
61+ def _compute_noise (gfun , x , p , sqrt_dt , z_t ):
62+ g = gfun (x , p )
63+ return g * sqrt_dt * z_t
6164
62- def theta (f , j , y0 , h , tf , * pars , th = 0.5 , sigma = 0.0 , tol = 1e-4 , key = None ):
65+
66+ def make_implicit_sde (dt , dfun , jfun , gfun , th = 0.5 , tol = 1e-4 , max_iters = 10 ):
6367 """
64- Implicit Theta integration scheme .
68+ Construct an implicit SDE integrator using the Theta method .
6569
6670 Parameters
6771 ----------
68- f : callable
69- Drift function f(y, *pars).
70- j : callable
71- Jacobian function j(y, *pars).
72- y0 : array
73- Initial state.
74- h : float
72+ dt : float
7573 Time step.
76- tf : float
77- Final time.
78- *pars : list
79- Parameters passed to f and j.
74+ dfun : callable
75+ Drift function dfun(x, p).
76+ jfun : callable
77+ Jacobian of drift function jfun(x, p).
78+ gfun : callable or float
79+ Diffusion function gfun(x, p) or constant sigma.
8080 th : float
81- Theta parameter (0.5 for trapezoidal/Crank-Nicolson, 1.0 for backward Euler).
82- sigma : float or array
83- Noise standard deviation.
81+ Theta parameter (0.5 for Crank-Nicolson, 1.0 for backward Euler).
8482 tol : float
8583 Tolerance for the Newton-Jacobi iteration.
86- key : jax.random.PRNGKey, optional
87- Random key for noise generation. If None, a default key is used .
84+ max_iters : int
85+ Maximum iterations for the Newton loop .
8886
8987 Returns
9088 -------
91- ts : array
92- Time points .
93- ys : array
94- Trajectory .
89+ step : callable
90+ Step function step(x, z_t, p) .
91+ loop : callable
92+ Loop function loop(x0, zs, p) .
9593 """
96- if key is None :
97- key = jax .random .PRNGKey (42 )
9894
99- num_steps = int (tf / h )
100- add_noise = np .any (sigma > 0.0 )
95+ if not hasattr (gfun , '__call__' ):
96+ sig = gfun
97+ gfun = lambda * _ : sig
10198
102- def scan_body (carry , _ ):
103- y0 , rng_key = carry
99+ sqrt_dt = np .sqrt (dt )
104100
105- f_y0 = (1 - th ) * f (y0 , * pars )
106- y1 = y0 + h * f_y0
101+ def step (x , z_t , p ):
102+ # Euler guess
103+ f_val = dfun (x , p )
104+ f_y0 = (1 - th ) * f_val
105+ y1_euler = x + dt * f_val # Initial guess using explicit Euler
107106
107+ # Refine if implicit
108108 def refine (y1 ):
109- dy = _theta_dy (f , j , h , th , y0 , y0 , f_y0 , pars , tol = tol )
110- y1 = y0 + dy
109+ # First Newton step
110+ # We use y1 (Euler guess) to compute Jacobian and Residual
111+ dy = _theta_dy (dfun , jfun , dt , th , x , y1 , f_y0 , p , tol = tol )
112+ y1 = y1 + dy
111113
112114 def refinement_cond (state ):
113115 y1 , dy , n_iter = state
114- return (n_iter < 10 ) & (np .linalg .norm (dy ) > tol )
116+ return (n_iter < max_iters ) & (np .linalg .norm (dy ) > tol )
115117
116118 def refinement_body (state ):
117119 y1 , _ , n_iter = state
118- dy = _theta_dy (f , j , h , th , y0 , y1 , f_y0 , pars , tol = tol )
120+ dy = _theta_dy (dfun , jfun , dt , th , x , y1 , f_y0 , p , tol = tol )
119121 y1 = y1 + dy
120122 return y1 , dy , n_iter + 1
121123
122124 init_state = (y1 , dy , 1 )
123125 final_state = jax .lax .while_loop (refinement_cond , refinement_body , init_state )
124126 return final_state [0 ]
125127
126- y1 = jax .lax .cond (th > 0.0 , refine , lambda x : x , y1 )
128+ y1 = jax .lax .cond (th > 0.0 , refine , lambda x : x , y1_euler )
127129
128- rng_key , step_key = jax .random .split (rng_key )
129- noise = jax .lax .cond (
130- add_noise ,
131- lambda k : np .sqrt (sigma ) * jax .random .normal (k , y0 .shape ),
132- lambda k : np .zeros_like (y0 ),
133- step_key
134- )
130+ noise = _compute_noise (gfun , x , p , sqrt_dt , z_t )
135131 y1 = y1 + noise
136132
137- return (y1 , rng_key ), y1
138-
139- _ , ys = jax .lax .scan (scan_body , (y0 , key ), None , length = num_steps )
133+ return y1
140134
141- # Prepend y0
142- ys = np .concatenate ([y0 [None , ...], ys ], axis = 0 )
143- ts = np .arange (num_steps + 1 ) * h
135+ @jax .jit
136+ def loop (x0 , zs , p ):
137+ def op (x , z ):
138+ x = step (x , z , p )
139+ return x , x
140+ _ , xs = jax .lax .scan (op , x0 , zs )
141+ return xs
144142
145- return ts , ys
143+ return step , loop
0 commit comments