Skip to content

Commit 7f72222

Browse files
Reflected mh (#3)
* need to commit before holidays * This should be the reflected mh implemented. * Finalised MMALA example
1 parent 86d6cc6 commit 7f72222

4 files changed

Lines changed: 364 additions & 79 deletions

File tree

Lines changed: 91 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -1,54 +1,119 @@
11
""" This algorithm implements Algorithm 4 of Maximal Couplings of the Metropolis-Hastings
22
Algorithm to compare with our method in the manifold MALA case. It is however coded in a fairly generic manner, so I
33
felt like putting it in the main code.
4+
The notations essentially follow from the paper notations but is all in log space.
45
"""
56
import jax
67
import jax.numpy as jnp
78

89
from coupled_rejection_sampling.utils import logsubexp
910

1011

11-
def reflection_coupled_mh(kernel, kernel_log_density):
12-
def r_xy_fn(x, y, x_prime):
13-
log_p_x_zp = kernel_log_density(x_prime, x)
14-
log_p_y_zp = kernel_log_density(x_prime, y)
15-
return logsubexp(log_p_x_zp, log_p_y_zp)
12+
def reflection_coupled_mh(proposal, proposal_log_density, target_log_density):
13+
log_f = _get_log_f(proposal_log_density, target_log_density)
14+
log_f_r = _get_log_f_r(log_f)
15+
log_f_t = _get_log_f_t(log_f)
1616

17-
def tr_xy_fn(x, y, x_prime):
18-
y_refl = _transport_fn(x, y, x_prime)
19-
r_xy_xp = r_xy_fn(x, y, x_prime)
20-
r_yx_yp = r_xy_fn(y, y, y_refl)
21-
return logsubexp(r_xy_xp, r_yx_yp)
17+
def log_acceptance_ratio(x, x_prime):
18+
log_q_x_x_prime = proposal_log_density(x, x_prime)
19+
log_q_x_prime_x = proposal_log_density(x_prime, x)
2220

21+
log_pi_x = target_log_density(x)
22+
log_pi_x_prime = target_log_density(x_prime)
23+
return jnp.minimum(0., log_pi_x_prime - log_pi_x + log_q_x_prime_x - log_q_x_x_prime)
2324

24-
def _residual(x, y):
25+
def mh_step(key, x):
26+
proposal_key, acceptance_key = jax.random.split(key)
27+
x_prime = proposal(proposal_key, x)
28+
log_alpha = log_acceptance_ratio(x, x_prime)
29+
log_u = jnp.log(jax.random.uniform(acceptance_key))
30+
return x_prime, log_u < log_alpha
31+
32+
def step_y(key, x, y, x_prime, accepted_x):
33+
y_prime = _transport_fn(x, y, x_prime)
34+
reflection_key, residual_key = jax.random.split(key)
35+
log_u = jnp.log(jax.random.uniform(reflection_key))
36+
cond = accepted_x & (log_u + log_f_r(x_prime, x, y) < log_f_r(y_prime, y, x))
37+
38+
return jax.lax.cond(cond, lambda _: (y_prime, True), lambda k: residual_y(k, y, x), residual_key)
39+
40+
def residual_y(key, y, x):
2541
def cond(carry):
2642
return ~carry[-1]
2743

2844
def body(carry):
29-
op_key, y_prop, _ = carry
30-
next_key, subkey1, subkey2 = jax.random.split(op_key)
31-
y_prop, accepted_y = kernel(subkey2, y)
45+
curr_key, *_ = carry
46+
next_key, mh_key, accept_key = jax.random.split(curr_key, 3)
47+
y_prime, accepted_y = mh_step(mh_key, y)
48+
log_u = jnp.log(jax.random.uniform(accept_key))
49+
stop = (~accepted_y) | (log_u + log_f(y, y_prime) < log_f_t(y_prime, y, x))
50+
return next_key, y_prime, accepted_y, stop
3251

33-
def if_accepted(_):
34-
log_u = jnp.log(jax.random.uniform(subkey1))
35-
log_p_y_yp = kernel_log_density(y_prop, y)
36-
log_p_res_y_yp = tr_xy_fn(y, x, y_prop)
37-
accept = log_u < log_p_res_y_yp - log_p_y_yp
38-
return y_prop, accept
39-
40-
return jax.lax.cond(accepted_y, lambda _: (y_prop, False), if_accepted, None)
52+
out = jax.lax.while_loop(cond, body, (key, y, False, False))
53+
return out[1], out[2]
4154

4255
def step(key, x, y):
43-
subkey1, subkey2, subkey3 = key
44-
x_prime, accepted_x = kernel(subkey1, x)
56+
x_key, coupling_key, residual_key = jax.random.split(key, 3)
57+
x_prime, accepted_x = mh_step(x_key, x)
58+
log_u = jnp.log(jax.random.uniform(coupling_key))
59+
next_x = jax.lax.select(accepted_x, x_prime, x)
60+
cond = accepted_x & (log_u + log_f(x, next_x) < log_f(y, next_x))
61+
62+
y_prime, accepted_y = jax.lax.cond(cond,
63+
lambda *_: (next_x, accepted_x),
64+
lambda k: step_y(k, x, y, next_x, accepted_x),
65+
residual_key)
66+
67+
next_y = jax.lax.select(accepted_y, y_prime, y)
68+
return (next_x, accepted_x), (next_y, accepted_y), cond
4569

70+
return step
4671

4772

4873
def _transport_fn(x, y, x_prime):
4974
r_curr = jnp.linalg.norm(y - x)
5075
e = jax.lax.select(r_curr < 1e-8, jnp.zeros_like(x), (y - x) / r_curr)
5176
x_diff = x_prime - x
5277
eta = x_diff - 2 * e * e.dot(x_diff)
53-
y = x + eta
54-
return y
78+
out = y + eta
79+
return out
80+
81+
82+
def _get_log_f(proposal_log_density, target_log_density):
83+
def log_f(x, x_prime):
84+
log_q_x_x_prime = proposal_log_density(x, x_prime)
85+
log_q_x_prime_x = proposal_log_density(x_prime, x)
86+
87+
log_pi_x = target_log_density(x)
88+
log_pi_x_prime = target_log_density(x_prime)
89+
90+
return jnp.minimum(log_q_x_x_prime, log_pi_x_prime + log_q_x_prime_x - log_pi_x)
91+
92+
return log_f
93+
94+
95+
def _get_log_f_m(log_f):
96+
def log_f_m(z, x, y):
97+
return jnp.minimum(log_f(x, z), log_f(y, z))
98+
99+
return log_f_m
100+
101+
102+
def _get_log_f_r(log_f):
103+
log_f_m = _get_log_f_m(log_f)
104+
105+
def log_f_r(x_prime, x, y):
106+
return logsubexp(log_f(x, x_prime), log_f_m(x_prime, x, y))
107+
108+
return log_f_r
109+
110+
111+
def _get_log_f_t(log_f):
112+
log_f_r = _get_log_f_r(log_f)
113+
114+
def log_f_t(y_prime, y, x):
115+
residual = log_f_r(y_prime, y, x)
116+
y_prime_t = _transport_fn(y, x, y_prime)
117+
return logsubexp(residual, jnp.minimum(residual, log_f_r(y_prime_t, x, y)))
118+
119+
return log_f_t

0 commit comments

Comments
 (0)