11""" This algorithm implements Algorithm 4 of Maximal Couplings of the Metropolis-Hastings
22Algorithm to compare with our method in the manifold MALA case. It is however coded in a fairly generic manner, so I
33felt like putting it in the main code.
4+ The notations essentially follow from the paper notations but is all in log space.
45"""
56import jax
67import jax .numpy as jnp
78
89from 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
4873def _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