Skip to content

Commit 3f6af15

Browse files
authored
Merge pull request #95 from ins-amu/jules/implicit-schemes-9047034567365458749
Add implicit Theta scheme and Jacobi solver
2 parents ec8c8cd + 45bec95 commit 3f6af15

3 files changed

Lines changed: 279 additions & 0 deletions

File tree

vbjax/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@ def _use_many_cores():
2121

2222
# import stuff
2323
from .loops import make_sde, make_ode, make_dde, make_sdde, heun_step, make_continuation
24+
from .implicit import jacobi, make_implicit_sde
2425
from .shtlc import make_shtdiff
2526
from .neural_mass import (
2627
JRState, JRTheta, jr_dfun, jr_default_theta,

vbjax/implicit.py

Lines changed: 143 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,143 @@
1+
import jax
2+
import jax.numpy as np
3+
4+
def jacobi(A, b, w=2.0/3.0, tol=1e-9, max_iters=100):
5+
"""
6+
Jacobi iterative linear solver for Ax = b.
7+
8+
Parameters
9+
----------
10+
A : array
11+
Matrix A.
12+
b : array
13+
Vector b.
14+
w : float
15+
Relaxation parameter.
16+
tol : float
17+
Tolerance.
18+
max_iters : int
19+
Maximum number of iterations.
20+
21+
Returns
22+
-------
23+
x : array
24+
Solution vector.
25+
n_iter : int
26+
Number of iterations performed.
27+
"""
28+
29+
# Precompute diagonal inverse and LU
30+
diag_A = np.diag(A)
31+
inv_diag_A = 1.0 / diag_A
32+
w_m_invD = w * np.diag(inv_diag_A)
33+
LU = A - np.diag(diag_A)
34+
35+
x0 = b
36+
37+
def cond(state):
38+
x, dx, n_iter = state
39+
return (n_iter < max_iters) & (np.linalg.norm(dx) > tol)
40+
41+
def body(state):
42+
x, dx, n_iter = state
43+
xn = w_m_invD @ (b - LU @ x) + (1.0 - w) * x
44+
dx = x - xn
45+
return xn, dx, n_iter + 1
46+
47+
dx = np.ones_like(x0)
48+
# Using jax.lax.while_loop for the iterative solver
49+
x, dx, n_iter = jax.lax.while_loop(cond, body, (x0, dx, 0))
50+
51+
return x, n_iter
52+
53+
54+
def _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
57+
58+
dy, _ = jacobi(J, F, w=2.0/3.0, tol=tol, max_iters=10)
59+
return dy
60+
61+
def _compute_noise(gfun, x, p, sqrt_dt, z_t):
62+
g = gfun(x, p)
63+
return g * sqrt_dt * z_t
64+
65+
66+
def make_implicit_sde(dt, dfun, jfun, gfun, th=0.5, tol=1e-4, max_iters=10):
67+
"""
68+
Construct an implicit SDE integrator using the Theta method.
69+
70+
Parameters
71+
----------
72+
dt : float
73+
Time step.
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.
80+
th : float
81+
Theta parameter (0.5 for Crank-Nicolson, 1.0 for backward Euler).
82+
tol : float
83+
Tolerance for the Newton-Jacobi iteration.
84+
max_iters : int
85+
Maximum iterations for the Newton loop.
86+
87+
Returns
88+
-------
89+
step : callable
90+
Step function step(x, z_t, p).
91+
loop : callable
92+
Loop function loop(x0, zs, p).
93+
"""
94+
95+
if not hasattr(gfun, '__call__'):
96+
sig = gfun
97+
gfun = lambda *_: sig
98+
99+
sqrt_dt = np.sqrt(dt)
100+
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
106+
107+
# Refine if implicit
108+
def refine(y1):
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
113+
114+
def refinement_cond(state):
115+
y1, dy, n_iter = state
116+
return (n_iter < max_iters) & (np.linalg.norm(dy) > tol)
117+
118+
def refinement_body(state):
119+
y1, _, n_iter = state
120+
dy = _theta_dy(dfun, jfun, dt, th, x, y1, f_y0, p, tol=tol)
121+
y1 = y1 + dy
122+
return y1, dy, n_iter + 1
123+
124+
init_state = (y1, dy, 1)
125+
final_state = jax.lax.while_loop(refinement_cond, refinement_body, init_state)
126+
return final_state[0]
127+
128+
y1 = jax.lax.cond(th > 0.0, refine, lambda x: x, y1_euler)
129+
130+
noise = _compute_noise(gfun, x, p, sqrt_dt, z_t)
131+
y1 = y1 + noise
132+
133+
return y1
134+
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
142+
143+
return step, loop

vbjax/tests/test_implicit.py

Lines changed: 135 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,135 @@
1+
2+
import jax
3+
import jax.numpy as np
4+
import vbjax
5+
import pytest
6+
7+
def test_jacobi():
8+
A = np.array([[4.0, 1.0], [1.0, 3.0]])
9+
b = np.array([1.0, 2.0])
10+
x, n_iter = vbjax.jacobi(A, b, tol=1e-6)
11+
expected = np.linalg.solve(A, b)
12+
assert np.allclose(x, expected, atol=1e-5)
13+
14+
def test_make_implicit_sde_linear():
15+
# dy = -k * y * dt
16+
k = 1.0
17+
def f(y, k):
18+
return -k * y
19+
20+
def j_fn(y, k):
21+
return -k * np.eye(y.size)
22+
23+
y0 = np.array([1.0])
24+
dt = 0.1
25+
tf = 1.0
26+
27+
# 0.5 = Crank-Nicolson -> trapezoidal rule
28+
step, loop = vbjax.make_implicit_sde(dt, f, j_fn, 0.0, th=0.5)
29+
30+
# Create noise (zeros since deterministic test)
31+
n_steps = int(tf / dt)
32+
zs = np.zeros((n_steps, 1))
33+
34+
ys = loop(y0, zs, k)
35+
36+
ts = np.arange(1, n_steps + 1) * dt
37+
exact = y0 * np.exp(-ts)
38+
39+
factor = (1 - k*dt/2) / (1 + k*dt/2)
40+
expected_numeric = y0 * (factor ** np.arange(1, n_steps + 1))
41+
42+
error = np.abs(ys.flatten() - expected_numeric)
43+
assert np.max(error) < 1e-5
44+
45+
# Also check against exact just to be sure it's close
46+
error_exact = np.abs(ys.flatten() - exact)
47+
assert np.max(error_exact) < 1e-2
48+
49+
50+
def test_make_implicit_sde_autodiff():
51+
# dy = -y^3
52+
def f(y, p):
53+
return -y**3
54+
55+
# Auto-diff Jacobian
56+
j_fn = jax.jacfwd(f)
57+
58+
y0 = np.array([1.0])
59+
dt = 0.1
60+
tf = 1.0
61+
62+
step, loop = vbjax.make_implicit_sde(dt, f, j_fn, 0.0, th=1.0) # Backward Euler
63+
64+
n_steps = int(tf / dt)
65+
zs = np.zeros((n_steps, 1))
66+
67+
ys = loop(y0, zs, None)
68+
69+
# Backward Euler: y_{n+1} = y_n - h * y_{n+1}^3
70+
# Check consistency
71+
for i in range(n_steps):
72+
prev = y0 if i == 0 else ys[i-1]
73+
curr = ys[i]
74+
res = curr + dt * curr**3 - prev
75+
assert np.abs(res) < 1e-4
76+
77+
@pytest.mark.benchmark(group="stiff_solver")
78+
def test_benchmark_heun_stiff(benchmark):
79+
k = 1000.0
80+
dt = 0.0001 # Stability limit requires small dt
81+
tf = 1.0
82+
n_steps = int(tf / dt)
83+
84+
def f(y, k):
85+
return -k * y
86+
87+
def g(y, k):
88+
return 0.1
89+
90+
y0 = np.ones(100)
91+
zs = jax.random.normal(jax.random.PRNGKey(0), (n_steps, 100))
92+
93+
_, loop = vbjax.make_sde(dt, f, g)
94+
95+
# Warmup / Compile
96+
loop(y0, zs, k).block_until_ready()
97+
98+
def run():
99+
return loop(y0, zs, k).block_until_ready()
100+
101+
benchmark(run)
102+
103+
@pytest.mark.benchmark(group="stiff_solver")
104+
def test_benchmark_implicit_stiff(benchmark):
105+
k = 1000.0
106+
dt = 0.01 # Implicit can take larger steps
107+
tf = 1.0
108+
n_steps = int(tf / dt)
109+
110+
def f(y, k):
111+
return -k * y
112+
113+
def j_fn(y, k):
114+
return -k * np.eye(y.size)
115+
116+
def g(y, k):
117+
return 0.1
118+
119+
y0 = np.ones(100)
120+
zs = jax.random.normal(jax.random.PRNGKey(0), (n_steps, 100))
121+
122+
step, loop = vbjax.make_implicit_sde(dt, f, j_fn, g, th=0.5)
123+
124+
# Warmup / Compile
125+
loop(y0, zs, k).block_until_ready()
126+
127+
def run():
128+
return loop(y0, zs, k).block_until_ready()
129+
130+
benchmark(run)
131+
132+
if __name__ == "__main__":
133+
test_jacobi()
134+
test_make_implicit_sde_linear()
135+
test_make_implicit_sde_autodiff()

0 commit comments

Comments
 (0)