%23%20%2F%2F%2F%20script%0A%23%20dependencies%20%3D%20%5B%0A%23%20%20%20%20%20%22jax%5Bcuda13%5D%3D%3D0.11.0%22%2C%0A%23%20%20%20%20%20%22marimo%22%2C%0A%23%20%20%20%20%20%22matplotlib%3D%3D3.11.0%22%2C%0A%23%20%20%20%20%20%22numpy%3D%3D2.5.1%22%2C%0A%23%20%20%20%20%20%22numpyro%3D%3D0.21.0%22%2C%0A%23%20%20%20%20%20%22pandas%3D%3D3.0.3%22%2C%0A%23%20%20%20%20%20%22python-dotenv%3D%3D1.2.2%22%2C%0A%23%20%20%20%20%20%22seaborn%3D%3D0.13.2%22%2C%0A%23%20%5D%0A%23%20requires-python%20%3D%20%22%3E%3D3.14%22%0A%23%20%2F%2F%2F%0A%0Aimport%20marimo%0A%0A__generated_with%20%3D%20%220.23.14%22%0Aapp%20%3D%20marimo.App(%0A%20%20%20%20width%3D%22medium%22%2C%0A%20%20%20%20app_title%3D%22Ancestral%20sampling%20with%20variational%20relaxation%22%2C%0A)%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_()%3A%0A%20%20%20%20import%20marimo%20as%20mo%0A%20%20%20%20from%20dotenv%20import%20load_dotenv%0A%20%20%20%20load_dotenv('.env')%0A%20%20%20%20from%20numpyro%20import%20distributions%20as%20dist%2C%20sample%2C%20factor%0A%20%20%20%20import%20jax%0A%20%20%20%20import%20jax.numpy%20as%20jnp%0A%20%20%20%20import%20time%0A%20%20%20%20from%20numpyro.infer%20import%20MCMC%2C%20NUTS%2C%20SVI%2C%20Trace_ELBO%2C%20Predictive%0A%20%20%20%20from%20numpyro.infer.autoguide%20import%20AutoDelta%2C%20AutoNormal%0A%20%20%20%20from%20numpyro.infer.initialization%20import%20init_to_value%0A%20%20%20%20import%20numpyro.optim%20as%20optim%0A%20%20%20%20import%20seaborn%20as%20sns%0A%20%20%20%20import%20numpy%20as%20np%0A%20%20%20%20import%20pandas%20as%20pd%0A%20%20%20%20from%20matplotlib%20import%20pyplot%20as%20plt%0A%20%20%20%20from%20jax.random%20import%20PRNGKey%0A%0A%20%20%20%20return%20(%0A%20%20%20%20%20%20%20%20AutoDelta%2C%0A%20%20%20%20%20%20%20%20AutoNormal%2C%0A%20%20%20%20%20%20%20%20MCMC%2C%0A%20%20%20%20%20%20%20%20NUTS%2C%0A%20%20%20%20%20%20%20%20PRNGKey%2C%0A%20%20%20%20%20%20%20%20Predictive%2C%0A%20%20%20%20%20%20%20%20SVI%2C%0A%20%20%20%20%20%20%20%20Trace_ELBO%2C%0A%20%20%20%20%20%20%20%20dist%2C%0A%20%20%20%20%20%20%20%20factor%2C%0A%20%20%20%20%20%20%20%20init_to_value%2C%0A%20%20%20%20%20%20%20%20jax%2C%0A%20%20%20%20%20%20%20%20jnp%2C%0A%20%20%20%20%20%20%20%20mo%2C%0A%20%20%20%20%20%20%20%20np%2C%0A%20%20%20%20%20%20%20%20optim%2C%0A%20%20%20%20%20%20%20%20pd%2C%0A%20%20%20%20%20%20%20%20plt%2C%0A%20%20%20%20%20%20%20%20sample%2C%0A%20%20%20%20%20%20%20%20sns%2C%0A%20%20%20%20%20%20%20%20time%2C%0A%20%20%20%20)%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(PRNGKey%2C%20mo%2C%20sns)%3A%0A%20%20%20%20KEY%20%3D%20PRNGKey(0)%0A%20%20%20%20N_PARTICLES%20%3D%201000%0A%20%20%20%20CONCENTRATION%20%3D%201.0%0A%20%20%20%20R%20%3D%200.010%0A%20%20%20%20STRENGTH%20%3D%20100.0%0A%0A%20%20%20%20sns.set_theme(%0A%20%20%20%20%20%20%20%20context%3D%22notebook%22%2C%0A%20%20%20%20%20%20%20%20style%3D%22ticks%22%2C%0A%20%20%20%20%20%20%20%20font%3D%22Inter%22%2C%0A%20%20%20%20%20%20%20%20rc%3D%7B%22svg.fonttype%22%3A%20%22none%22%2C%20%22savefig.format%22%3A%20%22svg%22%7D%2C%0A%20%20%20%20)%0A%0A%20%20%20%20def%20fig(f)%3A%0A%20%20%20%20%20%20%20%20return%20mo.as_html(f).style(%7B%22width%22%3A%20%22max-content%22%2C%20%22display%22%3A%20%22block%22%7D)%0A%0A%20%20%20%20return%20CONCENTRATION%2C%20KEY%2C%20N_PARTICLES%2C%20R%2C%20STRENGTH%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%20Ancestral%20sampling%20with%20variational%20relaxation%0A%20%20%20%20I%20recently%20discovered%20a%20little%20trick%20that%20worked%20surprisingly%20well%20for%20sampling%20from%20a%20complex%20prior%20distribution.%20Unsure%20how%20rigorous%20it%20is%20mathematically%20but%20thought%20it%20would%20be%20worth%20sharing%2Fdocumenting%20here.%20Here%20is%20how%20it%20goes.%0A%0A%20%20%20%20%23%23%20Easy%20and%20fast%20ancestral%20sampling%0A%20%20%20%20I%20used%20to%20think%20that%20some%20form%20of%20MCMC%20was%20always%20necessary%20to%20sample%20from%20any%20probabilistic%20program.%20This%20turns%20out%20to%20be%20incorrect%3B%20a%20lot%20of%20the%20time%20you%20can%20essentially%20step%20through%20the%20program%20and%20at%20every%20sample%20site%20simply%20do%20exactly%20that%3A%20sample%20from%20the%20respective%20distribution%2C%20potentially%20parameterised%20by%20previous%20variables.%20Take%20the%20example%20below%3A%20nothing%20is%20stopping%20us%20from%20sampling%20%60x%60%20from%20a%20beta%20distribution%2C%20then%20sampling%20%60y%60%20from%20another%20beta.%20It%20is%20clean%2C%20simple%20and%20much%20faster%20than%20MCMC%2C%20and%20we%20don't%20have%20to%20worry%20about%20step%20size%2C%20warm-up%2C%20effective%20sample%20size%2C%20etc.%20I%20have%20learned%20that%20this%20straightline%20%22execution%22%20mode%20is%20called%20ancestral%20sampling%2C%20presumably%20because%20there%20can%20even%20be%20dependencies%20between%20variables%20and%20their%20ancestors.%20As%20long%20as%20each%20variable%20is%20still%20something%20you%20know%20how%20to%20sample%20from%20by%20the%20time%20you%20reach%20it%20in%20the%20execution%20trace%2C%20ancestral%20sampling%20is%20possible.%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(CONCENTRATION%2C%20dist%2C%20sample)%3A%0A%20%20%20%20def%20random_points(n%2C%20scenes%3D1)%3A%0A%20%20%20%20%20%20%20%20x%20%3D%20sample('x'%2C%20dist.Beta(1.5%2C%201.5).expand((n%2C%20scenes)))%0A%20%20%20%20%20%20%20%20y%20%3D%20sample('y'%2C%20dist.Beta(CONCENTRATION%20%2B%20100.0%20*%20x%2C%20CONCENTRATION%20%2B%20100.0%20*%20x))%0A%0A%20%20%20%20return%20(random_points%2C)%0A%0A%0A%40app.cell%0Adef%20_(%0A%20%20%20%20KEY%2C%0A%20%20%20%20MCMC%2C%0A%20%20%20%20NUTS%2C%0A%20%20%20%20N_PARTICLES%2C%0A%20%20%20%20Predictive%2C%0A%20%20%20%20jax%2C%0A%20%20%20%20mo%2C%0A%20%20%20%20pd%2C%0A%20%20%20%20random_points%2C%0A%20%20%20%20time%2C%0A)%3A%0A%20%20%20%20%23%20Ancestral%20sampling%0A%20%20%20%20_t0%20%3D%20time.time()%0A%20%20%20%20p%20%3D%20Predictive(random_points%2C%20num_samples%3D1)%0A%20%20%20%20ancestral_samples%20%3D%20p(KEY%2C%20N_PARTICLES)%0A%20%20%20%20jax.block_until_ready(ancestral_samples)%0A%20%20%20%20ancestral_seconds%20%3D%20time.time()%20-%20_t0%0A%20%20%20%20ancestral_df%20%3D%20pd.DataFrame(%7Bk%3A%20v.ravel()%20for%20k%2C%20v%20in%20ancestral_samples.items()%7D)%0A%0A%20%20%20%20%23%20MCMC%20sampling%0A%20%20%20%20_t0%20%3D%20time.time()%0A%20%20%20%20mcmc%20%3D%20MCMC(NUTS(random_points)%2C%20num_warmup%3D500%2C%20num_samples%3D1%2C%20progress_bar%3DFalse)%0A%20%20%20%20mcmc.run(KEY%2C%20N_PARTICLES)%0A%20%20%20%20mcmc_samples%20%3D%20mcmc.get_samples()%0A%20%20%20%20jax.block_until_ready(mcmc_samples)%0A%20%20%20%20mcmc_seconds%20%3D%20time.time()%20-%20_t0%0A%20%20%20%20mcmc_df%20%3D%20pd.DataFrame(%7Bk%3A%20v.ravel()%20for%20k%2C%20v%20in%20mcmc_samples.items()%7D)%0A%0A%20%20%20%20mo.md(f%22Ancestral%20sampling%3A%20**%7Bancestral_seconds%3A.2f%7D%20s**%20%C2%B7%20MCMC%3A%20**%7Bmcmc_seconds%3A.2f%7D%20s**%22)%0A%20%20%20%20return%20ancestral_df%2C%20mcmc_df%0A%0A%0A%40app.cell%0Adef%20_(ancestral_df%2C%20mcmc_df%2C%20mo%2C%20plot_param_chooser%2C%20sns)%3A%0A%20%20%20%20a_fig%20%3D%20sns.jointplot(data%3Dancestral_df%2C%20x%3D'x'%2C%20y%3D'y'%2C%20**plot_param_chooser.value)%0A%20%20%20%20m_fig%20%3D%20sns.jointplot(data%3Dmcmc_df%2C%20x%3D'x'%2C%20y%3D'y'%2C%20color%3D'C1'%2C%20**plot_param_chooser.value)%0A%0A%20%20%20%20mo.vstack(%5B%0A%20%20%20%20%20%20%20%20plot_param_chooser%2C%0A%20%20%20%20%20%20%20%20mo.hstack(%5Bmo.vstack(%5Bmo.md(%22%23%23%23%23%20Ancestral%20sampling%22)%2C%20a_fig%5D%2C%20align%3D'center')%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20mo.vstack(%5Bmo.md(%22%23%23%23%23%20MCMC%22)%2C%20m_fig%5D%2C%20align%3D'center')%5D)%20%20%20%20%0A%20%20%20%20%5D%2C%20align%3D'center')%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(N_PARTICLES%2C%20mo)%3A%0A%20%20%20%20mo.md(rf%22%22%22%0A%20%20%20%20This%20is%20the%20sort%20of%20model%20that%20you%20might%20imagine%20using%20to%20represent%20a%20cohort%20of%20particles%20(here%20%60N_PARTICLES%60%20%3D%20%7BN_PARTICLES%7D).%20Ancestral%20sampling%20and%20NUTS%20are%20mostly%20on%20par%20in%20terms%20of%20speed%20and%20more%20or%20less%20portray%20the%20expected%20joint%20distributions.%0A%0A%20%20%20%20Now%2C%20let's%20say%20on%20top%20of%20this%20the%20particles%20don't%20like%20to%20be%20too%20close%20to%20each%20other%2C%20like%20molecules%20in%20a%20gas%20or%20people%20in%20a%20room.%20We%20might%20model%20this%20with%20a%20repulsive%20potential%20or%20%60factor%60%20in%20%60numpyro%60.%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(CONCENTRATION%2C%20R%2C%20STRENGTH%2C%20dist%2C%20factor%2C%20jnp%2C%20sample)%3A%0A%20%20%20%20def%20repulsive_points(n%2C%20scenes%3D1)%3A%0A%20%20%20%20%20%20%20%20x%20%3D%20sample('x'%2C%20dist.Beta(1.5%2C%201.5).expand((n%2C%20scenes)))%0A%20%20%20%20%20%20%20%20y%20%3D%20sample('y'%2C%20dist.Beta(CONCENTRATION%20%2B%20100.0%20*%20x%2C%20CONCENTRATION%20%2B%20100.0%20*%20x))%0A%20%20%20%20%20%20%20%20pos%20%3D%20jnp.stack(%5Bx%2C%20y%5D%2C%20axis%3D-1)%0A%20%20%20%20%20%20%20%20d%20%3D%20jnp.sqrt(((pos%5B%3A%2C%20None%5D%20-%20pos%5BNone%2C%20%3A%5D)%20**%202).sum(-1)%20%2B%201e-12)%0A%20%20%20%20%20%20%20%20overlap%20%3D%20jnp.clip(2%20*%20R%20-%20d%2C%20min%3D0.0)%20*%20(1%20-%20jnp.eye(n))%5B%3A%2C%20%3A%2C%20None%5D%0A%20%20%20%20%20%20%20%20factor('repulsion'%2C%20-STRENGTH%20*%200.5%20*%20overlap.sum())%0A%0A%20%20%20%20return%20(repulsive_points%2C)%0A%0A%0A%40app.cell%0Adef%20_(%0A%20%20%20%20KEY%2C%0A%20%20%20%20MCMC%2C%0A%20%20%20%20NUTS%2C%0A%20%20%20%20N_PARTICLES%2C%0A%20%20%20%20Predictive%2C%0A%20%20%20%20jax%2C%0A%20%20%20%20mo%2C%0A%20%20%20%20pd%2C%0A%20%20%20%20repulsive_points%2C%0A%20%20%20%20time%2C%0A)%3A%0A%20%20%20%20%23%20One%20configuration%2C%20sampled%20ancestrally%20(ignores%20the%20repulsion%20factor)%0A%20%20%20%20_t0%20%3D%20time.time()%0A%20%20%20%20rp%20%3D%20Predictive(repulsive_points%2C%20num_samples%3D100)%0A%20%20%20%20r_ancestral_samples%20%3D%20rp(KEY%2C%20N_PARTICLES)%0A%20%20%20%20jax.block_until_ready(r_ancestral_samples)%0A%20%20%20%20r_ancestral_seconds%20%3D%20time.time()%20-%20_t0%0A%20%20%20%20r_ancestral_df%20%3D%20pd.DataFrame(%7B'x'%3A%20r_ancestral_samples%5B'x'%5D%5B0%2C%20%3A%2C%200%5D%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20'y'%3A%20r_ancestral_samples%5B'y'%5D%5B0%2C%20%3A%2C%200%5D%7D)%0A%0A%20%20%20%20%23%20One%20configuration%20from%20MCMC%20(honours%20the%20repulsion%20factor)%0A%20%20%20%20_t0%20%3D%20time.time()%0A%20%20%20%20r_mcmc%20%3D%20MCMC(NUTS(repulsive_points)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20num_warmup%3D500%2C%20num_samples%3D100%2C%20progress_bar%3DFalse)%0A%20%20%20%20r_mcmc.run(KEY%2C%20N_PARTICLES)%0A%20%20%20%20r_mcmc_samples%20%3D%20r_mcmc.get_samples()%0A%20%20%20%20jax.block_until_ready(r_mcmc_samples)%0A%20%20%20%20r_mcmc_seconds%20%3D%20time.time()%20-%20_t0%0A%20%20%20%20r_mcmc_df%20%3D%20pd.DataFrame(%7B'x'%3A%20r_mcmc_samples%5B'x'%5D%5B0%2C%20%3A%2C%200%5D%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20'y'%3A%20r_mcmc_samples%5B'y'%5D%5B0%2C%20%3A%2C%200%5D%7D)%0A%0A%20%20%20%20mo.md(f%22Ancestral%20sampling%3A%20**%7Br_ancestral_seconds%3A.2f%7D%20s**%20%C2%B7%20MCMC%3A%20**%7Br_mcmc_seconds%3A.1f%7D%20s**%22)%0A%20%20%20%20return%20r_ancestral_samples%2C%20r_mcmc_samples%0A%0A%0A%40app.cell%0Adef%20_(%0A%20%20%20%20mo%2C%0A%20%20%20%20pd%2C%0A%20%20%20%20plot_param_chooser%2C%0A%20%20%20%20r_ancestral_samples%2C%0A%20%20%20%20r_mcmc_samples%2C%0A%20%20%20%20sample_slider%2C%0A%20%20%20%20sns%2C%0A)%3A%0A%20%20%20%20_r_sample%20%3D%20sample_slider.value%0A%20%20%20%20_r_ancestral_df%20%3D%20pd.DataFrame(%7B%0A%20%20%20%20%20%20%20%20%22x%22%3A%20r_ancestral_samples%5B%22x%22%5D%5B_r_sample%2C%20%3A%2C%200%5D%2C%0A%20%20%20%20%20%20%20%20%22y%22%3A%20r_ancestral_samples%5B%22y%22%5D%5B_r_sample%2C%20%3A%2C%200%5D%2C%0A%20%20%20%20%7D)%0A%20%20%20%20_r_mcmc_df%20%3D%20pd.DataFrame(%7B%0A%20%20%20%20%20%20%20%20%22x%22%3A%20r_mcmc_samples%5B%22x%22%5D%5B_r_sample%2C%20%3A%2C%200%5D%2C%0A%20%20%20%20%20%20%20%20%22y%22%3A%20r_mcmc_samples%5B%22y%22%5D%5B_r_sample%2C%20%3A%2C%200%5D%2C%0A%20%20%20%20%7D)%0A%20%20%20%20ra_fig%20%3D%20sns.jointplot(data%3D_r_ancestral_df%2C%20x%3D%22x%22%2C%20y%3D%22y%22%2C%20**plot_param_chooser.value)%0A%20%20%20%20rm_fig%20%3D%20sns.jointplot(data%3D_r_mcmc_df%2C%20x%3D%22x%22%2C%20y%3D%22y%22%2C%20color%3D%22C1%22%2C%20**plot_param_chooser.value)%0A%20%20%20%20ra_fig.ax_joint.set(xlim%3D(0%2C1)%2C%20ylim%3D(0%2C1))%0A%20%20%20%20rm_fig.ax_joint.set(xlim%3D(0%2C1)%2C%20ylim%3D(0%2C1))%0A%0A%20%20%20%20mo.vstack(%5B%0A%20%20%20%20%20%20%20%20mo.hstack(%5Bplot_param_chooser%2C%20sample_slider%5D)%2C%0A%20%20%20%20%20%20%20%20mo.hstack(%5Bmo.vstack(%5Bmo.md(%22%23%23%23%23%20Ancestral%20sampling%22)%2C%20ra_fig%5D%2C%20align%3D%22center%22)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20mo.vstack(%5Bmo.md(%22%23%23%23%23%20MCMC%22)%2C%20rm_fig%5D%2C%20align%3D%22center%22)%5D)%0A%20%20%20%20%5D%2C%20align%3D%22center%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20The%20added%20repulsion%20breaks%20our%20assumption%20about%20being%20able%20to%20sample%20from%20each%20site%20sequentially.%20Now%20that%20we%20are%20sampling%20from%20an%20arbitrary%20joint%20distribution%2C%20ancestral%20sampling%20using%20%60Predictive%60%20simply%20ignores%20the%20%60factor%60%2C%20producing%20an%20identical%20result%20to%20before.%20I%20had%20to%20fiddle%20with%20MCMC%20settting%20to%20make%20it%20work%20within%20a%20reasonable%20amount%20of%20time%2C%20but%20the%20end%20result%20shows%20that%20it%20is%20clearly%20working%2C%20even%20if%20painfully%20slow.%20There%20are%20some%20concerning%20signs%3A%20note%20how%20the%20marginal%20of%20%24x%24%20is%20spread%20to%20the%20left.%20This%20is%20not%20unexpected%20but%20as%20the%20case%20often%20is%20with%20MCMC%2C%20it%20has%20the%20smell%20of%20something%20that%20needs%20a%20bit%20more%20poking%20before%20I%20trust%20it.%20What%20if%20we%20could%20use%20ancestral%20sampling%20to%20generate%20lots%20of%20proposal%20initial%20distributions%20that%20could%20then%20be%20%22relaxed%22%20to%20reflect%20the%20repulsive%20potential.%20Let's%20see%20...%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(%0A%20%20%20%20KEY%2C%0A%20%20%20%20N_PARTICLES%2C%0A%20%20%20%20Predictive%2C%0A%20%20%20%20SVI%2C%0A%20%20%20%20Trace_ELBO%2C%0A%20%20%20%20guide_chooser%2C%0A%20%20%20%20init_to_value%2C%0A%20%20%20%20jax%2C%0A%20%20%20%20mo%2C%0A%20%20%20%20optim%2C%0A%20%20%20%20repulsive_points%2C%0A%20%20%20%20time%2C%0A)%3A%0A%20%20%20%20guide_type%2C%20iterations%20%3D%20guide_chooser.value%0A%20%20%20%20_t0%20%3D%20time.time()%0A%20%20%20%20relax_init%20%3D%20Predictive(repulsive_points%2C%20num_samples%3D1)(KEY%2C%20N_PARTICLES%2C%20scenes%3D100)%0A%20%20%20%20relax_init_x%2C%20relax_init_y%20%3D%20relax_init%5B'x'%5D%5B0%5D%2C%20relax_init%5B'y'%5D%5B0%5D%0A%20%20%20%20relax_guide%20%3D%20guide_type(%0A%20%20%20%20%20%20%20%20repulsive_points%2C%0A%20%20%20%20%20%20%20%20init_loc_fn%3Dinit_to_value(values%3D%7B'x'%3A%20relax_init_x%2C%20'y'%3A%20relax_init_y%7D)%2C%0A%20%20%20%20)%0A%20%20%20%20relax_svi%20%3D%20SVI(repulsive_points%2C%20relax_guide%2C%20optim.Adam(1e-4)%2C%20Trace_ELBO())%0A%20%20%20%20relax_result%20%3D%20relax_svi.run(KEY%2C%20iterations%2C%20N_PARTICLES%2C%20scenes%3D100%2C%20progress_bar%3DFalse)%0A%20%20%20%20relaxed_posterior%20%3D%20relax_guide.sample_posterior(KEY%2C%20relax_result.params%2C%20N_PARTICLES%2C%20100%2C%20sample_shape%3D(1%2C))%0A%20%20%20%20relaxed_x%2C%20relaxed_y%20%3D%20relaxed_posterior%5B'x'%5D.squeeze()%2C%20relaxed_posterior%5B'y'%5D.squeeze()%0A%0A%20%20%20%20jax.block_until_ready((relaxed_x%2C%20relaxed_y))%0A%20%20%20%20relax_seconds%20%3D%20time.time()%20-%20_t0%0A%0A%20%20%20%20mo.md(f%22Ancestral%20%2B%20variational%20relaxation%20of%20**100**%20configurations%3A%20%22%0A%20%20%20%20%20%20%20%20%20%20f%22**%7Brelax_seconds%3A.1f%7D%20s**%20(~%7Brelax_seconds%20%2F%20100%3A.2f%7D%20s%20per%20configuration)%22)%0A%20%20%20%20return%20relax_guide%2C%20relax_init_x%2C%20relax_init_y%2C%20relaxed_x%2C%20relaxed_y%0A%0A%0A%40app.cell%0Adef%20_(%0A%20%20%20%20guide_chooser%2C%0A%20%20%20%20mo%2C%0A%20%20%20%20pd%2C%0A%20%20%20%20plot_param_chooser%2C%0A%20%20%20%20relax_init_x%2C%0A%20%20%20%20relax_init_y%2C%0A%20%20%20%20relaxed_x%2C%0A%20%20%20%20relaxed_y%2C%0A%20%20%20%20scene_slider%2C%0A%20%20%20%20sns%2C%0A)%3A%0A%20%20%20%20_s%20%3D%20scene_slider.value%0A%20%20%20%20proposal_df%20%3D%20pd.DataFrame(%7B'x'%3A%20relax_init_x%5B%3A%2C%20_s%5D%2C%20'y'%3A%20relax_init_y%5B%3A%2C%20_s%5D%7D)%0A%20%20%20%20relaxed_df%20%3D%20pd.DataFrame(%7B'x'%3A%20relaxed_x%5B%3A%2C%20_s%5D%2C%20'y'%3A%20relaxed_y%5B%3A%2C%20_s%5D%7D)%0A%20%20%20%20pj_fig%20%3D%20sns.jointplot(data%3Dproposal_df%2C%20x%3D'x'%2C%20y%3D'y'%2C%20**plot_param_chooser.value)%0A%20%20%20%20rv_fig%20%3D%20sns.jointplot(data%3Drelaxed_df%2C%20x%3D'x'%2C%20y%3D'y'%2C%20color%3D'C2'%2C%20**plot_param_chooser.value)%0A%20%20%20%20pj_fig.ax_joint.set(xlim%3D(0%2C1)%2C%20ylim%3D(0%2C1))%0A%20%20%20%20rv_fig.ax_joint.set(xlim%3D(0%2C1)%2C%20ylim%3D(0%2C1))%0A%0A%20%20%20%20mo.vstack(%5B%0A%20%20%20%20%20%20%20%20mo.hstack(%5Bguide_chooser%2C%20plot_param_chooser%2C%20scene_slider%5D)%2C%0A%20%20%20%20%20%20%20%20mo.hstack(%5Bmo.vstack(%5Bmo.md(%22%23%23%23%23%20Ancestral%20proposal%22)%2C%20pj_fig%5D%2C%20align%3D'center')%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20mo.vstack(%5Bmo.md(%22%23%23%23%23%20Relaxed%22)%2C%20rv_fig%5D%2C%20align%3D'center')%5D)%0A%20%20%20%20%5D%2C%20align%3D'center')%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20Much%20better!%20Still%20squeezed%20in%20the%20middle%20at%20high%20%24x%24%2C%20where%20the%20prior%20dominates%20but%20nicely%20spaced%20out%20in%20the%20flat%20middle%2C%20as%20expected.%20Quantitative%20comparison%20of%20the%20target%20log%20density%20shows%20much%20higher%20likelihoods%20than%20the%20ancestral%20starting%20point%20and%20potentially%20even%20MCMC.%20Note%20this%20is%20not%20necessarily%20a%20good%20thing%3B%20the%20relaxation%20may%20be%20(and%20indeed%20is)%20pushing%20the%20distribution%20towards%20the%20maximum%20a%20posteriori%20(MAP)%20point.%20In%20this%20case%2C%20the%20fact%20that%20we%20deliberately%20ran%20a%20very%20small%20number%20of%20steps%20and%20learning%20rate%2C%20_i.e._%20aborted%20SVI%20early%20was%20precisely%20to%20prevent%20this.%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(mcmc_logp%2C%20mo%2C%20plt%2C%20proposal_logp%2C%20relaxed_logp%2C%20sns)%3A%0A%20%20%20%20_f%2C%20_ax%20%3D%20plt.subplots(figsize%3D(6.5%2C%203.6))%0A%20%20%20%20sns.histplot(proposal_logp%2C%20color%3D'C0'%2C%20label%3D'Ancestral%20proposals'%2C%20kde%3DTrue%2C%20stat%3D'count'%2C%20ax%3D_ax)%0A%20%20%20%20sns.histplot(mcmc_logp%2C%20color%3D'C1'%2C%20label%3D'MCMC'%2C%20kde%3DTrue%2C%20stat%3D'count'%2C%20ax%3D_ax)%0A%20%20%20%20sns.histplot(relaxed_logp%2C%20color%3D'C2'%2C%20label%3D'Relaxed%20ancestral'%2C%20kde%3DTrue%2C%20stat%3D'count'%2C%20ax%3D_ax)%0A%20%20%20%20_ax.set(xlabel%3D'Target%20log-density'%2C%20ylabel%3D'%23%20samples')%0A%20%20%20%20_ax.legend(frameon%3DFalse)%0A%20%20%20%20sns.despine(_f)%0A%0A%20%20%20%20mo.vstack(%5B%0A%20%20%20%20%20%20%20%20mo.md(%22%23%23%23%20Unnormalised%20target%20log%20density%22)%2C%0A%20%20%20%20%20%20%20%20_f%2C%0A%20%20%20%20%5D%2C%20align%3D'center')%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%23%20Other%20potential%20applications%0A%20%20%20%20Some%20problems%20don't%20_look_%20like%20%60factor%60%20but%20are%20effectively%20the%20same%20thing.%20Any%20observed%20variables%2C%20for%20instance%2C%20including%20between%20latent%20sites.%20I've%20recently%20realised%20that%20a%20generative%20process%20that%20I%20wrote%20about%20a%20couple%20of%20years%20ago%20%5Bhere%5D(https%3A%2F%2Fhessammehr.github.io%2Fblog%2Fposts%2F2024-12-26-locally-constrained.html)%20has%20a%20name%3A%20a%20%5BMarkov%20random%20field%5D(https%3A%2F%2Fen.wikipedia.org%2Fwiki%2FMarkov_random_field).%20This%20type%20of%20process%20is%20essentially%20equivalent%20to%20making%20adding%20observations%20between%20latent%20variables%2C%20for%20example%20positing%20that%20%24x_n%20-%20x_%7Bn-1%7D%20%5Csim%20%5Cmathcal%7BN%7D(%5Cmu%2C%20%5Csigma)%24%20for%20certain%20or%20all%20%24n%24.%20Again%2C%20we're%20adding%20a%20potential%2Ffactor%20to%20the%20joint%20probability%20distribution.%0A%0A%20%20%20%20Anyway%20this%20was%20a%20lot%20of%20fun%20to%20discover%2C%20and%20I%20have%20yet%20to%20find%20out%20whether%20it%20is%20commonly%20used%20or%20known%20about.%20Hopefully%20interesting%2Fuseful!%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%23%20Update%3A%20More%20on%20caveats%0A%20%20%20%20There%20are%20some%20caveats%20to%20this%20method%20that%20I'm%20aware%20of%2C%20and%20some%20that%20I%20am%20not!%20The%20main%20one%20as%20noted%20above%20is%20that%20SVI%20%2B%20the%20%60AutoDelta%60%20guide%20alone%20will%20simply%20squeeze%20the%20distribution%20towards%20MAP.%20I%20think%20we're%20doing%20quite%20well%20here%2C%20but%20that's%20because%20I%20chose%20a%20number%20of%20SVI%20steps%20that%20was%20just%20enough%20to%20relax%20the%20added%20factor%20without%20significant%20drift%20towards%20MAP.%20Too%20many%20steps%20may%20be%20detectable%20as%20the%20joint%20log-density%20distribution%20of%20samples%20starts%20to%20collapse%20to%20a%20spike%2C%20like%20in%20the%20plot%20below.%20I%20actually%20had%20to%20update%20the%20notebook%20to%20work%20with%20%24%5Ctext%7BBeta%7D(%5Ccdot%2C%20%5Ccdot)%24%20parameters%20%24%3E1%24%20because%20between%200%20and%201%20the%20beta%20distribution%20has%20singularities%20as%200%20and%2For%201%20which%20complicate%20the%20story.%0A%0A%20%20%20%20I%20have%20briefly%20experimented%20with%20using%20%60AutoNormal%60%20instead%20of%20%60AutoDelta%60%2C%20where%20the%20former's%20entropy%20term%20will%20presumably%20prevent%20this%20type%20of%20collapse%20when%20maximising%20the%20evidence%20lower%20bound%20(ELBO).%20In%20fact%2C%20you%20can%20switch%20the%20guide%20to%20%60AutoNormal%60%20above%20and%20see%20how%20it%20affects%20the%20results.%20There%20is%20definitely%20no%20collapse%20even%20after%20a%20large%20number%20of%20SVI%20steps%2C%20though%20from%20a%20quick%20visual%20check%20I'm%20not%20sure%20overall%20whether%20the%20goal%20of%20particles%20avoiding%20each%20other%20is%20fulfilled%20in%20this%20case.%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(compute_svi_logp_path%2C%20mo%2C%20plt%2C%20sns)%3A%0A%20%20%20%20_logp_fig%2C%20_logp_ax%20%3D%20plt.subplots(figsize%3D(7.2%2C%204.4))%0A%0A%20%20%20%20svi_logp_path%20%3D%20compute_svi_logp_path((0%2C%20100%2C%20300%2C%201000%2C%2010000))%0A%0A%20%20%20%20sns.histplot(%0A%20%20%20%20%20%20%20%20data%3Dsvi_logp_path%2C%0A%20%20%20%20%20%20%20%20x%3D%22Target%20log%20density%22%2C%0A%20%20%20%20%20%20%20%20hue%3D%22SVI%20steps%22%2C%0A%20%20%20%20%20%20%20%20bins%3D50%2C%0A%20%20%20%20%20%20%20%20stat%3D%22density%22%2C%0A%20%20%20%20%20%20%20%20common_norm%3DFalse%2C%0A%20%20%20%20%20%20%20%20common_bins%3DTrue%2C%0A%20%20%20%20%20%20%20%20linewidth%3D1.5%2C%0A%20%20%20%20%20%20%20%20ax%3D_logp_ax%2C%0A%20%20%20%20)%0A%20%20%20%20_logp_ax.set_ylabel(%22Density%22)%0A%20%20%20%20sns.despine(_logp_fig)%0A%0A%20%20%20%20mo.vstack(%5B%0A%20%20%20%20%20%20%20%20mo.md(%22%23%23%23%20Joint%20log%20density%20across%20100%20configurations%22)%2C%0A%20%20%20%20%20%20%20%20_logp_fig%2C%0A%20%20%20%20%5D%2C%20align%3D'center')%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20plot_param_options%20%3D%20%7B%0A%20%20%20%20%20%20%20%20'Scatter'%3A%20dict(s%3D10%2C%20alpha%3D0.25%2C%20marginal_kws%3D%7B'stat'%3A%20'density'%7D)%2C%0A%20%20%20%20%20%20%20%20'KDE'%3A%20dict(kind%3D'kde'%2C%20fill%3DTrue%2C%20levels%3D20%2C%20bw_method%3D0.25%2C%20marginal_kws%3D%7B'clip'%3A%20(0%2C%201)%2C%20'cut'%3A%200%7D)%0A%20%20%20%20%7D%0A%0A%20%20%20%20plot_param_chooser%20%3D%20mo.ui.dropdown(plot_param_options%2C%20value%3D'Scatter'%2C%20label%3D'Plot%20type')%0A%20%20%20%20scene_slider%20%3D%20mo.ui.slider(0%2C%2099%2C%20value%3D0%2C%20label%3D'Configuration')%0A%20%20%20%20sample_slider%20%3D%20mo.ui.slider(0%2C%2099%2C%20value%3D0%2C%20label%3D'Sample')%0A%20%20%20%20return%20plot_param_chooser%2C%20sample_slider%2C%20scene_slider%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(%0A%20%20%20%20jax%2C%0A%20%20%20%20np%2C%0A%20%20%20%20r_mcmc_samples%2C%0A%20%20%20%20relax_init_x%2C%0A%20%20%20%20relax_init_y%2C%0A%20%20%20%20relaxed_x%2C%0A%20%20%20%20relaxed_y%2C%0A%20%20%20%20repulsive_points%2C%0A)%3A%0A%20%20%20%20from%20numpyro.infer.util%20import%20log_density%0A%0A%20%20%20%20def%20config_logp(xc%2C%20yc)%3A%0A%20%20%20%20%20%20%20%20n%20%3D%20xc.shape%5B0%5D%0A%20%20%20%20%20%20%20%20ld%2C%20_%20%3D%20log_density(repulsive_points%2C%20(n%2C)%2C%20%7B'scenes'%3A%201%7D%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%7B'x'%3A%20xc%5B%3A%2C%20None%5D%2C%20'y'%3A%20yc%5B%3A%2C%20None%5D%7D)%0A%20%20%20%20%20%20%20%20return%20ld%0A%0A%20%20%20%20_logp_batch%20%3D%20jax.jit(jax.vmap(config_logp%2C%20in_axes%3D(1%2C%201)))%0A%20%20%20%20relaxed_logp%20%3D%20np.asarray(_logp_batch(relaxed_x%2C%20relaxed_y))%0A%20%20%20%20proposal_logp%20%3D%20np.asarray(_logp_batch(relax_init_x%2C%20relax_init_y))%0A%20%20%20%20mcmc_logp%20%3D%20np.asarray(_logp_batch(%0A%20%20%20%20%20%20%20%20r_mcmc_samples%5B'x'%5D%5B...%2C%200%5D.T%2C%0A%20%20%20%20%20%20%20%20r_mcmc_samples%5B'y'%5D%5B...%2C%200%5D.T%2C%0A%20%20%20%20))%0A%20%20%20%20return%20config_logp%2C%20mcmc_logp%2C%20proposal_logp%2C%20relaxed_logp%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(%0A%20%20%20%20KEY%2C%0A%20%20%20%20N_PARTICLES%2C%0A%20%20%20%20SVI%2C%0A%20%20%20%20Trace_ELBO%2C%0A%20%20%20%20config_logp%2C%0A%20%20%20%20jax%2C%0A%20%20%20%20np%2C%0A%20%20%20%20optim%2C%0A%20%20%20%20pd%2C%0A%20%20%20%20relax_guide%2C%0A%20%20%20%20repulsive_points%2C%0A)%3A%0A%20%20%20%20def%20compute_svi_logp_path(checkpoints)%3A%0A%20%20%20%20%20%20%20%20relax_svi%20%3D%20SVI(repulsive_points%2C%20relax_guide%2C%20optim.Adam(3e-5)%2C%20Trace_ELBO())%0A%20%20%20%20%20%20%20%20_state%20%3D%20relax_svi.init(KEY%2C%20N_PARTICLES%2C%20scenes%3D100)%0A%20%20%20%20%20%20%20%20_batched_logp%20%3D%20jax.jit(jax.vmap(config_logp%2C%20in_axes%3D(1%2C%201)))%0A%0A%20%20%20%20%20%20%20%20%40jax.jit%0A%20%20%20%20%20%20%20%20def%20_advance(_current_state%2C%20_count)%3A%0A%20%20%20%20%20%20%20%20%20%20%20%20return%20jax.lax.fori_loop(%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%200%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20_count%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20lambda%20_%2C%20_inner_state%3A%20relax_svi.update(%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20_inner_state%2C%20N_PARTICLES%2C%20scenes%3D100%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20)%5B0%5D%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20_current_state%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20)%0A%0A%20%20%20%20%20%20%20%20_frames%20%3D%20%5B%5D%0A%20%20%20%20%20%20%20%20_previous%20%3D%200%0A%20%20%20%20%20%20%20%20for%20_checkpoint%20in%20checkpoints%3A%0A%20%20%20%20%20%20%20%20%20%20%20%20_state%20%3D%20_advance(_state%2C%20_checkpoint%20-%20_previous)%0A%20%20%20%20%20%20%20%20%20%20%20%20_params%20%3D%20relax_svi.get_params(_state)%0A%20%20%20%20%20%20%20%20%20%20%20%20relaxed_posterior%20%3D%20relax_guide.sample_posterior(KEY%2C%20_params%2C%20N_PARTICLES%2C%20100%2C%20sample_shape%3D(1%2C))%0A%20%20%20%20%20%20%20%20%20%20%20%20_values%20%3D%20np.asarray(%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20_batched_logp(relaxed_posterior%5B'x'%5D.squeeze()%2C%20relaxed_posterior%5B'y'%5D.squeeze())%0A%20%20%20%20%20%20%20%20%20%20%20%20)%0A%20%20%20%20%20%20%20%20%20%20%20%20_frames.append(%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20pd.DataFrame(%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%7B%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22SVI%20steps%22%3A%20str(_checkpoint)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Target%20log%20density%22%3A%20_values%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%7D%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20)%0A%20%20%20%20%20%20%20%20%20%20%20%20)%0A%20%20%20%20%20%20%20%20%20%20%20%20_previous%20%3D%20_checkpoint%0A%20%20%20%20%20%20%20%20return%20pd.concat(_frames%2C%20ignore_index%3DTrue)%0A%0A%20%20%20%20return%20(compute_svi_logp_path%2C)%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(AutoDelta%2C%20AutoNormal%2C%20mo)%3A%0A%20%20%20%20guide_options%20%3D%20%7B%0A%20%20%20%20%20%20%20%20'Delta'%3A%20(AutoDelta%2C%20300)%2C%0A%20%20%20%20%20%20%20%20'Normal'%3A%20(AutoNormal%2C%2010000)%0A%20%20%20%20%7D%0A%0A%20%20%20%20guide_chooser%20%3D%20mo.ui.dropdown(guide_options%2C%20value%3D'Delta'%2C%20label%3D'Guide%20type')%0A%20%20%20%20return%20(guide_chooser%2C)%0A%0A%0Aif%20__name__%20%3D%3D%20%22__main__%22%3A%0A%20%20%20%20app.run()%0A
1c04e47ab814171cde9acddbeb6a7413