Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
Stefan Ivanovic 1 Ge Liu 1 Mohammed El-Kebir 1 2
arXiv:2606.07400v1 [cs.LG] 5 Jun 2026
Abstract
mating unknown true states that are often combinatorial in nature, such as chemical reaction pathways (Ji & Deng, 2021), transportation networks (Biagioni & Eriksson, 2012), evolutionary trees (Felsenstein, 1981), molecular conformations (Noé et al., 2013), gene regulatory networks (Friedman et al., 2000), etc. Importantly, we do not observe the latent states Si and are instead given indirect observations Xi . These problems are self-supervised in the sense that the statistical process Pr(X | S) linking each indirect observation Xi to each latent state Si is at least partially understood, rather than relying purely on unsupervised reconstruction of the observations Xi . Modeling mechanistic latent states can be formulated as optimizing a generative modelQ Pr(S P | θ) with parameters θ that maximize the probability i S Pr(Xi | S) Pr(S | θ) of our observations Xi . The resulting inference problem is often solved using classical approaches such as expectation maximization (Dempster et al., 1977). However, classical approaches often struggle to scale to combinatorially large latent state spaces.
Many scientific problems require inferring unobserved mechanistic latent states from indirect observations. While classical approaches, including expectation maximization, do not scale to combinatorially large spaces, deep learning approaches such as variational autoencoders typically form artificial latent states rather than reconstructing the mechanistic ground-truth states. Here, we introduce GReinSS, a policy learning framework that uses dynamically rescaled rewards to learn latent state distributions that maximize the observed data likelihood. We show that GReinSS accurately reconstructs simulated latent sets and latent graphs, outperforming alternative policy learning and generative modeling baselines. Additionally, GReinSS reconstructs isoforms from real short-read RNA sequencing data that better match isoforms detected by orthogonal long-read sequencing than the standard RSEM algorithm. Overall, GReinSS is a principled and practically effective approach for generative modeling and inference of combinatorial latent states from indirect observations.
Reinforcement learning (RL) naturally fits combinatorial problems by sequentially generating their underlying structured latent states. While RL has been used to maximize expected rewards (Sutton et al., 1998), maximize entropy (Haarnoja et al., 2017), and match trajectory distributions to reward distributions (Malkin et al., 2022), and policy gradients have been used within variational inference to optimize surrogate objectives (Mnih & Rezende, 2016; Mnih & Gregor, 2014), these approaches do not directly optimize the marginal likelihood of indirect observations.
1. Introduction Learning latent states that explain observed data is a central goal of machine learning and probabilistic modeling. General-purpose unsupervised approaches such as clustering (MacQueen, 1965), topic modeling (Hofmann, 1999), and representation learning (Bengio et al., 2013) find artificial latent states that represent the variation in observed data Xi without attempting to match some unknown true state Si∗ . On the other hand, practical and/or scientific problems often require inferring mechanistic latent states approxi-
Our contribution: Here, we introduce Generative Reinforcement Learning of Structured States (GReinSS), a framework for optimizing the distributionQPr(S | θ) of states S to maximize the overall probability i Pr(Xi | θ) of observations Xi . GReinSS uses policy gradients with dynamically scaled rewards to ensure that the policy’s parameters θ are Q updated in the direction of maximizing the probability i Pr(Xi | θ) of the observations Xi . In contrast to standard RL formulations, where policy gradients optimize expected return under a stationary reward function, GReinSS uses RL machinery as an optimization tool for probabilistic modeling over discrete latent states.
1
Siebel School of Computing and Data Science, University of Illinois at Urbana-Champaign, IL 61801, USA 2 Cancer Center at Illinois, University of Illinois Urbana-Champaign, IL 61801, USA. Correspondence to: Mohammed El-Kebir <[email protected]>.
We compare GReinSS with existing approaches on both simulated latent graph inference, simulated latent set re-
Proceedings of the 43 rd International Conference on Machine Learning, Seoul, South Korea. PMLR 306, 2026. Copyright 2026 by the author(s).
1
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
a
b input: X1:N <latexit sha1_base64="dB9y+I/li4qlGBEs5909fUOdJuU=">AAAB7nicbVDLSgNBEOyNrxhfUY9eBoPgKWSDRPEU9OJJIpgHJEuYnfQmQ2Znl5lZISz5CC8eFPHq93jzb5wke9DEgoaiqpvuLj8WXJtK5dvJra1vbG7ltws7u3v7B8XDo5aOEsWwySIRqY5PNQousWm4EdiJFdLQF9j2x7czv/2ESvNIPppJjF5Ih5IHnFFjpXann7rX99N+sVQpV+Ygq8TNSAkyNPrFr94gYkmI0jBBte661dh4KVWGM4HTQi/RGFM2pkPsWippiNpL5+dOyZlVBiSIlC1pyFz9PZHSUOtJ6NvOkJqRXvZm4n9eNzHBlZdyGScGJVssChJBTERmv5MBV8iMmFhCmeL2VsJGVFFmbEIFG4K7/PIqaVXLbq1ce7go1W+yOPJwAqdwDi5cQh3uoAFNYDCGZ3iFNyd2Xpx352PRmnOymWP4A+fzB8DpjzU=</latexit>
S <latexit sha1_base64="DXccWvWphvaYCjWBT/AqJvJwhWw=">AAAB8nicbVDLSgMxFM34rPVVdekmWARXZaZIdVl047KifcB0KJk004ZmkiG5I5Shn+HGhSJu/Rp3/o2ZdhbaeiBwOOdecu4JE8ENuO63s7a+sbm1Xdop7+7tHxxWjo47RqWasjZVQuleSAwTXLI2cBCsl2hG4lCwbji5zf3uE9OGK/kI04QFMRlJHnFKwEp+PyYwpkRkD7NBperW3DnwKvEKUkUFWoPKV3+oaBozCVQQY3yvnkCQEQ2cCjYr91PDEkInZMR8SyWJmQmyeeQZPrfKEEdK2ycBz9XfGxmJjZnGoZ3MI5plLxf/8/wUousg4zJJgUm6+ChKBQaF8/vxkGtGQUwtIVRzmxXTMdGEgm2pbEvwlk9eJZ16zWvUGveX1eZNUUcJnaIzdIE8dIWa6A61UBtRpNAzekVvDjgvzrvzsRhdc4qdE/QHzucPkbqRdg==</latexit>
state: Si⇤: <latexit sha1_base64="6gV5f0BP07mN1yN7q2Xc3RrGWpc=">AAAB7nicbVDLSgNBEOz1GeMr6tHLYBDEQ8gGiR6DXjxGNA9I1jA7mU2GzM4uM71CWPIRXjwo4tXv8ebfOEn2oIkFDUVVN91dfiyFwXL521lZXVvf2Mxt5bd3dvf2CweHTRMlmvEGi2Sk2z41XArFGyhQ8nasOQ19yVv+6Gbqt564NiJSDziOuRfSgRKBYBSt1Lp/PO+lYtIrFMul8gxkmbgZKUKGeq/w1e1HLAm5QiapMR23EqOXUo2CST7JdxPDY8pGdMA7lioacuOls3Mn5NQqfRJE2pZCMlN/T6Q0NGYc+rYzpDg0i95U/M/rJBhcealQcYJcsfmiIJEEIzL9nfSF5gzl2BLKtLC3EjakmjK0CeVtCO7iy8ukWSm51VL17qJYu87iyMExnMAZuHAJNbiFOjSAwQie4RXenNh5cd6dj3nripPNHMEfOJ8/DgSPaA==</latexit>
output: ✓ <latexit sha1_base64="eCnCl7MWU2tTtly26VfwdxByNiY=">AAAB7XicbVBNS8NAEN3Ur1q/qh69LBbBU0mKVI9FLx4r2FpoQ9lsJ+3azSbsToQS+h+8eFDEq//Hm//GbZuDtj4YeLw3w8y8IJHCoOt+O4W19Y3NreJ2aWd3b/+gfHjUNnGqObR4LGPdCZgBKRS0UKCETqKBRYGEh2B8M/MfnkAbEat7nCTgR2yoRCg4Qyu1ezgCZP1yxa26c9BV4uWkQnI0++Wv3iDmaQQKuWTGdL1agn7GNAouYVrqpQYSxsdsCF1LFYvA+Nn82ik9s8qAhrG2pZDO1d8TGYuMmUSB7YwYjsyyNxP/87ophld+JlSSIii+WBSmkmJMZ6/TgdDAUU4sYVwLeyvlI6YZRxtQyYbgLb+8Stq1qlev1u8uKo3rPI4iOSGn5Jx45JI0yC1pkhbh5JE8k1fy5sTOi/PufCxaC04+c0z+wPn8AaoJjzU=</latexit>
<latexit sha1_base64="kDphGC8QqPt+N7/2YOdMCtyjVms=">AAACEHicbVC7TgJBFJ31ifhatbSZSIxgQVhi0JJoY4lBHgm7ktlhgAmzj8zcNZLNfoKNv2JjoTG2lnb+jQNsoeBJbnJyzr259x43FFxBqfRtLC2vrK6tZzaym1vbO7vm3n5TBZGkrEEDEci2SxQT3GcN4CBYO5SMeK5gLXd0NfFb90wqHvi3MA6Z45GBz/ucEtBS1zyxazJfx7bHe9iGIQNSwDYJQxk84Fh7yV18muTrha6ZKxVLU+BFYqUkh1LUuuaX3Qto5DEfqCBKdaxyCE5MJHAqWJK1I8VCQkdkwDqa+sRjyomnDyX4WCs93A+kLh/wVP09ERNPqbHn6k6PwFDNexPxP68TQf/CibkfRsB8OlvUjwSGAE/SwT0uGQUx1oRQyfWtmA6JJBR0hlkdgjX/8iJplotWpVi5OctVL9M4MugQHaE8stA5qqJrVEMNRNEjekav6M14Ml6Md+Nj1rpkpDMH6A+Mzx+cwJu4</latexit>
Xi : Si Xi i 2 [N ] <latexit sha1_base64="wGjaJcVIOLgboMBiCoYZvyawOsk=">AAAB6nicbVBNS8NAEJ34WetX1aOXxSJ4KkmR6rHoxWNF+wFtKJvtpF262YTdjVBCf4IXD4p49Rd589+4bXPQ1gcDj/dmmJkXJIJr47rfztr6xubWdmGnuLu3f3BYOjpu6ThVDJssFrHqBFSj4BKbhhuBnUQhjQKB7WB8O/PbT6g0j+WjmSToR3QoecgZNVZ66PR5v1R2K+4cZJV4OSlDjka/9NUbxCyNUBomqNZdr5oYP6PKcCZwWuylGhPKxnSIXUsljVD72fzUKTm3yoCEsbIlDZmrvycyGmk9iQLbGVEz0sveTPzP66YmvPYzLpPUoGSLRWEqiInJ7G8y4AqZERNLKFPc3krYiCrKjE2naEPwll9eJa1qxatVaveX5fpNHkcBTuEMLsCDK6jDHTSgCQyG8Ayv8OYI58V5dz4WrWtOPnMCf+B8/gA14o3F</latexit>
<latexit sha1_base64="eCnCl7MWU2tTtly26VfwdxByNiY=">AAAB7XicbVBNS8NAEN3Ur1q/qh69LBbBU0mKVI9FLx4r2FpoQ9lsJ+3azSbsToQS+h+8eFDEq//Hm//GbZuDtj4YeLw3w8y8IJHCoOt+O4W19Y3NreJ2aWd3b/+gfHjUNnGqObR4LGPdCZgBKRS0UKCETqKBRYGEh2B8M/MfnkAbEat7nCTgR2yoRCg4Qyu1ezgCZP1yxa26c9BV4uWkQnI0++Wv3iDmaQQKuWTGdL1agn7GNAouYVrqpQYSxsdsCF1LFYvA+Nn82ik9s8qAhrG2pZDO1d8TGYuMmUSB7YwYjsyyNxP/87ophld+JlSSIii+WBSmkmJMZ6/TgdDAUU4sYVwLeyvlI6YZRxtQyYbgLb+8Stq1qlev1u8uKo3rPI4iOSGn5Jx45JI0yC1pkhbh5JE8k1fy5sTOi/PufCxaC04+c0z+wPn8AaoJjzU=</latexit>
<latexit sha1_base64="pE2tow+YrUQ9saCg+e8v4o/lavI=">AAAB8nicbVBNS8NAEN34WetX1aOXxSJ4kJIUqR6LXjxWsB+QhrLZTtqlm03YnQil9Gd48aCIV3+NN/+N2zYHbX0w8Hhvhpl5YSqFQdf9dtbWNza3tgs7xd29/YPD0tFxyySZ5tDkiUx0J2QGpFDQRIESOqkGFocS2uHobua3n0AbkahHHKcQxGygRCQ4Qyv5nZ64pF0cArJeqexW3DnoKvFyUiY5Gr3SV7ef8CwGhVwyY3yvmmIwYRoFlzAtdjMDKeMjNgDfUsViMMFkfvKUnlulT6NE21JI5+rviQmLjRnHoe2MGQ7NsjcT//P8DKObYCJUmiEovlgUZZJiQmf/077QwFGOLWFcC3sr5UOmGUebUtGG4C2/vEpa1YpXq9Qersr12zyOAjklZ+SCeOSa1Mk9aZAm4SQhz+SVvDnovDjvzseidc3JZ07IHzifP5ajkNM=</latexit>
ing parameters θ given latent variables S (M-step) (Dempster et al., 1977; Wu, 1983). This can be highly effective but value PNrequires the ability to compute the expectation (t+1) ))] of the i=1 ESi ∼Pr(S|Xi ,θ (t) ) [log(Pr(Xi , Si | θ complete-data log-likelihood. In settings where the latent state space is exponentially large, calculating this expectation is generally intractable. A classical exception is hidden Markov models, where the exponentially large state space has a Markov structure that allows exact expectation computation via dynamic programming (Rabiner, 1989). In Section 3.3, we describe how GEM can be adapted to use modern machine learning methods, including variational autoencoders, autoregressive models, and discrete diffusion, to approximately compute and optimize this expectation value on general exponentially large state spaces. More recently, variational inference (VI) (Wainwright et al., 2008) has emerged as a principled framework for learning probabilistic models on exponentially large latent spaces. Rather than directly optimizing the marginal likelihood Pr(X1:N | θ) using a known distribution Pr(X | S), VI solves a related problem in which the objective is replaced by an ELBO surrogate objective that includes an explicit variational posterior model qϕ (S | X) parameterized by ϕ.
<latexit sha1_base64="9/ke/u62nC+DnZA1ZlSFTSjgoeI=">AAAB8HicbVDLTgJBEOzFF+IL9ehlIjHxRFhi0CPRi0eM8jCwIbPDLEyYnd3M9JqQDV/hxYPGePVzvPk3DrAHBSvppFLVne4uP5bCYKXy7eTW1jc2t/LbhZ3dvf2D4uFRy0SJZrzJIhnpjk8Nl0LxJgqUvBNrTkNf8rY/vpn57SeujYjUA05i7oV0qEQgGEUrPfZGFNP7aV/0i6VKuTIHWSVuRkqQodEvfvUGEUtCrpBJakzXrcbopVSjYJJPC73E8JiyMR3yrqWKhtx46fzgKTmzyoAEkbalkMzV3xMpDY2ZhL7tDCmOzLI3E//zugkGV14qVJwgV2yxKEgkwYjMvicDoTlDObGEMi3srYSNqKYMbUYFG4K7/PIqaVXLbq1cu7so1a+zOPJwAqdwDi5cQh1uoQFNYBDCM7zCm6OdF+fd+Vi05pxs5hj+wPn8Af1NkI0=</latexit>
output: Ŝi
Pr(S | ✓) ⇡ Pr⇤ (S) <latexit sha1_base64="kDphGC8QqPt+N7/2YOdMCtyjVms=">AAACEHicbVC7TgJBFJ31ifhatbSZSIxgQVhi0JJoY4lBHgm7ktlhgAmzj8zcNZLNfoKNv2JjoTG2lnb+jQNsoeBJbnJyzr259x43FFxBqfRtLC2vrK6tZzaym1vbO7vm3n5TBZGkrEEDEci2SxQT3GcN4CBYO5SMeK5gLXd0NfFb90wqHvi3MA6Z45GBz/ucEtBS1zyxazJfx7bHe9iGIQNSwDYJQxk84Fh7yV18muTrha6ZKxVLU+BFYqUkh1LUuuaX3Qto5DEfqCBKdaxyCE5MJHAqWJK1I8VCQkdkwDqa+sRjyomnDyX4WCs93A+kLh/wVP09ERNPqbHn6k6PwFDNexPxP68TQf/CibkfRsB8OlvUjwSGAE/SwT0uGQUx1oRQyfWtmA6JJBR0hlkdgjX/8iJplotWpVi5OctVL9M4MugQHaE8stA5qqJrVEMNRNEjekav6M14Ml6Md+Nj1rpkpDMH6A+Mzx+cwJu4</latexit>
<latexit sha1_base64="tSbfGwJQT0JAzt8DtHddR/MHIAo=">AAACAHicbVC7TsMwFHV4lvIKMDCwWFRIiKFKKlQYK1gYi0ofUhMix3Vaq45j2Q6iirLwKywMIMTKZ7DxN7iPAVqOdKWjc+7VvfeEglGlHefbWlpeWV1bL2wUN7e2d3btvf2WSlKJSRMnLJGdECnCKCdNTTUjHSEJikNG2uHweuy3H4hUNOF3eiSIH6M+pxHFSBspsA+9AdJZIw8o9JAQMnmEjfuzgAZ2ySk7E8BF4s5ICcxQD+wvr5fgNCZcY4aU6roVof0MSU0xI3nRSxURCA9Rn3QN5Sgmys8mD+TwxCg9GCXSFNdwov6eyFCs1CgOTWeM9EDNe2PxP6+b6ujSzygXqSYcTxdFKYM6geM0YI9KgjUbGYKwpOZWiAdIIqxNZkUTgjv/8iJpVcputVy9PS/VrmZxFMAROAanwAUXoAZuQB00AQY5eAav4M16sl6sd+tj2rpkzWYOwB9Ynz8ohpYj</latexit>
<latexit sha1_base64="tSbfGwJQT0JAzt8DtHddR/MHIAo=">AAACAHicbVC7TsMwFHV4lvIKMDCwWFRIiKFKKlQYK1gYi0ofUhMix3Vaq45j2Q6iirLwKywMIMTKZ7DxN7iPAVqOdKWjc+7VvfeEglGlHefbWlpeWV1bL2wUN7e2d3btvf2WSlKJSRMnLJGdECnCKCdNTTUjHSEJikNG2uHweuy3H4hUNOF3eiSIH6M+pxHFSBspsA+9AdJZIw8o9JAQMnmEjfuzgAZ2ySk7E8BF4s5ICcxQD+wvr5fgNCZcY4aU6roVof0MSU0xI3nRSxURCA9Rn3QN5Sgmys8mD+TwxCg9GCXSFNdwov6eyFCs1CgOTWeM9EDNe2PxP6+b6ujSzygXqSYcTxdFKYM6geM0YI9KgjUbGYKwpOZWiAdIIqxNZkUTgjv/8iJpVcputVy9PS/VrmZxFMAROAanwAUXoAZuQB00AQY5eAav4M16sl6sd+tj2rpkzWYOwB9Ynz8ohpYj</latexit>
Ŝ ⇡ S ⇤ Ŝii ⇡ Sii⇤ <latexit sha1_base64="tSbfGwJQT0JAzt8DtHddR/MHIAo=">AAACAHicbVC7TsMwFHV4lvIKMDCwWFRIiKFKKlQYK1gYi0ofUhMix3Vaq45j2Q6iirLwKywMIMTKZ7DxN7iPAVqOdKWjc+7VvfeEglGlHefbWlpeWV1bL2wUN7e2d3btvf2WSlKJSRMnLJGdECnCKCdNTTUjHSEJikNG2uHweuy3H4hUNOF3eiSIH6M+pxHFSBspsA+9AdJZIw8o9JAQMnmEjfuzgAZ2ySk7E8BF4s5ICcxQD+wvr5fgNCZcY4aU6roVof0MSU0xI3nRSxURCA9Rn3QN5Sgmys8mD+TwxCg9GCXSFNdwov6eyFCs1CgOTWeM9EDNe2PxP6+b6ujSzygXqSYcTxdFKYM6geM0YI9KgjUbGYKwpOZWiAdIIqxNZkUTgjv/8iJpVcputVy9PS/VrmZxFMAROAanwAUXoAZuQB00AQY5eAav4M16sl6sd+tj2rpkzWYOwB9Ynz8ohpYj</latexit>
<latexit sha1_base64="Q3I6jKAhzjhukMe1weFv2AUmIIk=">AAAB6nicbVDLTgJBEOzFF+IL9ehlIjHxRFhi0CPRi0cM8khgQ2aHWZgwO7uZ6TUhGz7BiweN8eoXefNvHGAPClbSSaWqO91dfiyFwUrl28ltbG5t7+R3C3v7B4dHxeOTtokSzXiLRTLSXZ8aLoXiLRQoeTfWnIa+5B1/cjf3O09cGxGpR5zG3AvpSIlAMIpWajYHYlAsVcqVBcg6cTNSggyNQfGrP4xYEnKFTFJjem41Ri+lGgWTfFboJ4bHlE3oiPcsVTTkxksXp87IhVWGJIi0LYVkof6eSGlozDT0bWdIcWxWvbn4n9dLMLjxUqHiBLliy0VBIglGZP43GQrNGcqpJZRpYW8lbEw1ZWjTKdgQ3NWX10m7WnZr5drDVal+m8WRhzM4h0tw4RrqcA8NaAGDETzDK7w50nlx3p2PZWvOyWZO4Q+czx8uRI3A</latexit>
✓
⇤
Pr(S | ✓) ⇡ Pr (S) observation:
c input: Xi , ✓
<latexit sha1_base64="wGjaJcVIOLgboMBiCoYZvyawOsk=">AAAB6nicbVBNS8NAEJ34WetX1aOXxSJ4KkmR6rHoxWNF+wFtKJvtpF262YTdjVBCf4IXD4p49Rd589+4bXPQ1gcDj/dmmJkXJIJr47rfztr6xubWdmGnuLu3f3BYOjpu6ThVDJssFrHqBFSj4BKbhhuBnUQhjQKB7WB8O/PbT6g0j+WjmSToR3QoecgZNVZ66PR5v1R2K+4cZJV4OSlDjka/9NUbxCyNUBomqNZdr5oYP6PKcCZwWuylGhPKxnSIXUsljVD72fzUKTm3yoCEsbIlDZmrvycyGmk9iQLbGVEz0sveTPzP66YmvPYzLpPUoGSLRWEqiInJ7G8y4AqZERNLKFPc3krYiCrKjE2naEPwll9eJa1qxatVaveX5fpNHkcBTuEMLsCDK6jDHTSgCQyG8Ayv8OYI58V5dz4WrWtOPnMCf+B8/gA14o3F</latexit>
<latexit sha1_base64="Ab/FTAqRhqvv9RJDbDk/vx6YizY=">AAAB8HicbVBNSwMxEJ2tX7V+VT16CRbBU9ktUj0WvXiSCvZDtkvJptk2NMkuSVYoS3+FFw+KePXnePPfmLZ70NYHA4/3ZpiZFyacaeO6305hbX1jc6u4XdrZ3ds/KB8etXWcKkJbJOax6oZYU84kbRlmOO0mimIRctoJxzczv/NElWaxfDCThAYCDyWLGMHGSo8M9ZhE/l3QL1fcqjsHWiVeTiqQo9kvf/UGMUkFlYZwrLXv1RITZFgZRjidlnqppgkmYzykvqUSC6qDbH7wFJ1ZZYCiWNmSBs3V3xMZFlpPRGg7BTYjvezNxP88PzXRVZAxmaSGSrJYFKUcmRjNvkcDpigxfGIJJorZWxEZYYWJsRmVbAje8surpF2revVq/f6i0rjO4yjCCZzCOXhwCQ24hSa0gICAZ3iFN0c5L86787FoLTj5zDH8gfP5A8gPj8M=</latexit>
states S <latexit sha1_base64="73DcjFo2H9mWC0jR7+y3HvuFJ74=">AAAB6HicbVDLTgJBEOz1ifhCPXqZSEw8kV1i0CPRi0eI8khgQ2aHXhiZnd3MzJoQwhd48aAxXv0kb/6NA+xBwUo6qVR1p7srSATXxnW/nbX1jc2t7dxOfndv/+CwcHTc1HGqGDZYLGLVDqhGwSU2DDcC24lCGgUCW8Hodua3nlBpHssHM07Qj+hA8pAzaqxUv+8Vim7JnYOsEi8jRchQ6xW+uv2YpRFKwwTVuuOVE+NPqDKcCZzmu6nGhLIRHWDHUkkj1P5kfuiUnFulT8JY2ZKGzNXfExMaaT2OAtsZUTPUy95M/M/rpCa89idcJqlByRaLwlQQE5PZ16TPFTIjxpZQpri9lbAhVZQZm03ehuAtv7xKmuWSVylV6pfF6k0WRw5O4QwuwIMrqMId1KABDBCe4RXenEfnxXl3Phata042cwJ/4Hz+ALQ3jOQ=</latexit>
Figure 1. Overview of GReinSS. a Many scientific problems ∗ consist of estimating latent states S1∗ , . . . SN ∈ S only given indirect observations X1:N = X1 , . . . , XN . These latent states are sampled from an unobserved underlying distribution Pr∗ (S), which we model as Pr(S | θ) using parameters θ. This yields two problems. b First, a learning problem of identifying parameters θ that maximize the data likelihood Pr(X1:N | θ). c Second, an inference problem of estimating Si∗ with Ŝi = argmax Pr(S | Xi ) Pr(S | θ) given Xi and θ. GReinSS solves both problems using policy gradients with dynamic rewards.
construction, and the real biological problem of RNA isoform reconstruction from short-read sequencing data. On simulations, GReinSS reliably reconstructs ground-truth latent states with higher accuracy than standard policy gradients, GFlowNets, local search, and generalized expectation maximization implemented with variational autoencoders, diffusion models, and autoregressive models. On the real short-read RNA sequencing data from GTEx (Lonsdale et al., 2013), GReinSS outperforms the standard RSEM (Li & Dewey, 2011) method used by GTEx in terms of predicting gene isoforms validated by additional long-read RNA sequencing. These results demonstrate that GReinSS successfully extends modern generative modeling to the inference of discrete latent states from indirect observation data through a policy learning formulation.
Beyond EM and VI, reinforcement learning has been widely used for generating structured data using discrete action spaces (Yu et al., 2017). For instance, GFlowNets are used to learn policies that generate states S with probabilities Pr(S | θ) proportional to their reward (Malkin et al., 2022). However, this method assumes known rewards for terminal states S rather than optimizing the distribution of terminal states to maximize the probability Pr(X1 , . . . , XN | θ) of indirect observation data. Many approaches have used adaptive, normalized, or shaped reward functions; however, these are typically used to encourage exploration or stabilize training rather than to match some latent distribution of states (Ibrahim et al., 2024; Yuan et al., 2023). GReinSS instead constructs adaptive rewards that enable policy gradients to learn a distribution Pr(S | θ) over latent states S that maximizes the likelihood Pr(X1 , . . . , XN | θ) of indirect observation data X1 , . . . , XN . While the aforementioned methods do not generally solve the same problem as GReinSS, we identify several special cases where they do (Section 3.3). Moreover, we use these methods as baselines in our evaluation (Section 4.1).
1.1. Related Works There are several related works aimed at explaining observation data X1 , . . . , XN through latent states S. The simplest approach for inferring a latent state S from an observation Xi given the probability function Pr(Xi | S) is directly optimizing argmaxS Pr(Xi | S) via local search. This, however, does not leverage information across observations Xi through a shared model Pr(S | θ) of latent state generation. Variational autoencoders effectively learn a distribution over artificial latent states that maximizes the probability of generating observation data X1 , . . . , XN (Kingma & Welling, 2013). However, these artificial latent states belong to a completely separate vector space distinct from the mechanistic latent states that underlie the scientific problem.
In addition to related general methods, two previous papers have utilized the underlying technique of GReinSS without generalizing to describe the full GReinSS procedure (Ivanovic & El-Kebir, 2023; 2025) (see Section B.4 for details). We build upon these early applications and (i) formulate general probabilistic learning and inference problems on arbitrary discrete structures (Section 2); (ii) provide a theoretically grounded solution to these problems by establishing a connection between policy-gradient training and
The classical approaches of expectation maximization (EM) and generalized expectation maximization (GEM) for inferring problem-specific mechanistic latent states/variables S operate by alternating between inference over latent variables S given parameters θ (E-step) followed by updat2
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
maximum likelihood estimation in discrete combinatorial spaces (Section 3.1); (iii) formulate a theory for optimal off-policy sampling (Section 3.2); (iv) relate existing ML approaches to GReinSS (Section 3.3); and (v) provide ablations and comparisons with existing ML methods (Section 4). Conflict of Interest Disclosure:
instance, if multiple pieces of data are generated from each state Si∗ , one can group these sub-observations into one full observation Xi (see Section B.1 and Section 4.1). Additionally, multiple environments v, each with (i) their own ∗ ∗ distributions Pr∗v (S) of latent states, (ii) lists Sv,1 , . . . Sv,N v of ground-truth latent states, (iii) distributions Prv (X | S) v of observations, and (iv) lists X1v , . . . XN v of observations, can be subsumed in the original problem definitions (see Section B.1 and Section 4.2).
None.
2. Problem Definition
3. Methods
Let Pr∗ (S) be the ground-truth distribution on the set S of latent states S, which can represent graphs, sequences, sets, positions in space, or any discrete data structure that can be generated via policy learning. We are given a sample of N items from this distribution. Importantly, we do not directly ∗ ∗ observe the latent states S1:N = S1∗ , . . . , SN ; instead we are only given indirect observations X1:N = X1 , . . . , XN . In addition, we are given the ability to compute the probability Pr(X | S) of any observation X given any latent state S. As such, our observations X1:N are generated from ∗ Pr(X | S1∗ ), . . . , Pr(X | SN ). We aim to approximate ∗ Pr (S) via a generative model Pr(S | θ) of latent states S with parameters θ (Figure 1a). We note that we place no restrictions on the generative model, which can be a variational autoencoder, discrete diffusion, autoregressive generation, policy-based generation, etc. Therefore, the probability of an observation X given parameters θ is X Pr(X | θ) = Pr(X | S) Pr(S | θ). (1)
Proofs for theorems in this section are provided in Section A. 3.1. The GReinSS Method We solve the learning problem (Problem 2.1) of identifying parameters θ maximizing Pr(X1:N | θ) via gradient d descent by estimating the gradient dθ log(Pr(X1:N | θ)) of the data log-likelihood. In the following, we show how to accomplish this using policy gradients and a dynamically changing reward function. We define θ as the parameters of our policy. The policy defines a probability distribution Pr(τ | θ) over trajectories τ , which are sequences of actions concluding with a terminal state S(τ ). In our case, the terminal state S(τ ) of each trajectory τ corresponds to a latent state in S. As such, we define Pr(X | τ ) = Pr(X | S(τ )). Analogous to Equation (1), we define Pr(X | θ) = Eτ [Pr(X | τ )] = Eτ [Pr(X | S(τ ))]
S∈S
(2)
where τ is drawn from the distribution Pr(τ | θ). In practice, we estimate the quantity Pr(Xi | θ) by sampling, i.e., PM 1 Pr(Xi | θ) ≈ M j=1 Pr(Xi | τj ) where τ1 , . . . τM are sampled from Pr(τ | θ). This leads us to the following important theorem. d Theorem 3.1. The policy gradient Eτ [r(τ ) dθ log(Pr(τ | θ))] with dynamically changing rewards
The overall probability Pr(X1:N | θ) of the observations X1:N given our model parameters θ thus equals QN P i=1 S∈S Pr(Xi | S) Pr(S | θ). We pose the following learning problem to train Pr(S | θ) towards matching Pr∗ (S) without directly observing Pr∗ (S) (Figure 1b). Problem 2.1 (L ATENT S TATE M ODELING FROM I NDI RECT O BSERVATIONS). Given indirect observations X1:N and the probability function Pr(X | S), find model parameters θ that maximize the probability Pr(X1:N | θ) of the observations.
r(τ ) =
N X Pr(Xi | τ ) i=1
Pr(Xi | θ)
(3)
d is an unbiased estimator of the gradient dθ log(Pr(X1:N | θ)) of the log-likelihood objective.
Next, we pose the following inference problem to estimate ∗ the unobserved ground-truth latent states S1∗ , . . . , SN using the learned parameters θ (Figure 1c).
Intuitively, the denominator Pr(Xi | θ) rescales each observation’s contribution to the total reward such that trajectories are rewarded based on their proportional contribution to Pr(Xi | θ) rather than the raw probability Pr(Xi | τ ). This rescaling results in solving for the optimal distribution over trajectories Pr(τ | θ) rather than converging to one highest reward trajectory as illustrated in Section B.2 and Figure S1. As such, we can apply standard policy gradient updates to solve Problem 2.1.
Problem 2.2 (L ATENT S TATE I NFERENCE). Given indirect observations X1:N , the probability function Pr(X | S), and model parameters θ, for each observation Xi find the latent state Ŝi ∈ S that maximizes the probability Pr(Ŝi | Xi , θ) = Pr(Xi | Ŝi ) Pr(Ŝi | θ). Problems 2.1 and 2.2 also apply in more complex situations with distinct groups of observations and states. For 3
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients Table 1. Comparison of GReinSS to baseline methods. While local search can infer individual latent states Si from each observation Xi , it does not utilize parameters θ. While variational autoencoders (VAE), autoregression, and diffusion can be adapted to solve Problem 2.1 via Generalized Expectation-Maximization (GEM), it requires an inexact modification of the GEM algorithm. On the other hand, the two policy learning (PL) baselines maximize proxy objectives.
Corollary 3.2. GReinSS’s application of standard policy gradient updates to θ using the reward r(τ ) performs gradient ascent on the data log-likelihood log(Pr(X1:N | θ)). In other words, after each gradient update yielding parameters θ, GReinSS updates the dynamic rewards r(τ ) using these new parameters via (2) and (3), which in turn are used for the next gradient update of θ. We note that in standard d RL, policy gradients Eτ [r(τ ) dθ log(Pr(τ | θ))] are used to maximize the expectation value Eτ [r(τ )] of the rewards, given a fixed reward function r(τ ) independent of the policy parameters θ. Our dynamic rewards, which depend on parameters θ, allow the same policy gradients procedure to instead optimize the data log-likelihood log(Pr(X1:N | θ)) of our observations X1:N . Importantly, despite the rewards d being calculated using θ, the gradient dθ is only applied to log(Pr(τ | θ)) not r(τ ).
METHOD
PARADIGM
MAXIMIZES OBJ .
Pr(X1:N | θ) LOCAL SEARCH
SEARCH
VAE
GEM GEM GEM PL PL PL
AUTOREGRESSION DIFFUSION NAIVE POLICY GRAD . GF LOW N ETS GR EIN SS
NO APPROX . APPROX . APPROX . NO NO YES
τ . Thus, the off-policy sampling procedure only allows trajectories that generate a set of nodes corresponding to some observation Xi in X1:N . As another example, CNRein (Ivanovic & El-Kebir, 2025) first uses a simple nonmachine learning algorithm called CNNaive to give initial estimates of plausible latent states and then biases trajectories towards these plausible latent states. After applying off-policy learning, one can account for this in policy gradients using importance sampling (Section B.6).
After optimizing Pr(X1:N | θ) to solve Problem 2.1, one can solve Problem 2.2 by sampling latent states and approximating which latent state Ŝi maximizes Pr(Ŝi | θ) Pr(Xi | Ŝi ) for each data point i (details in Section B.5). Although Theorem 3.1 calculates the reward over all observations, policy gradients still give an unbiased estimator of the log-likelihood objective if one applies minibatching to calculate rewards (Theorem A.1). Specifically, for a P mini-batch B ⊂ [N ] the mini-batch reward is rB (τ ) = i∈B Pr(Xi | τ )/ Pr(Xi | θ). Analogous to mini-batch stochastic gradient ascent in standard likelihoodbased learning, applying mini-batching to GReinSS can enable scaling to very large datasets.
3.3. Baselines and Special Cases Maximum likelihood generative modeling special case: We start with the following special case, where the learning problem simplifies to maximum likelihood generative modeling with the log-likelihood objective.
3.2. Off-policy Learning and Importance Sampling
Lemma 3.4. Let the given observations equal the groundtruth states, i.e., Xi = Si∗ such that Pr(Xi | S) = 1 if S = Xi = Si∗ and 0 otherwise. Then, Problem 2.1 of solving argmaxθ Pr(X1:N | θ) simplifies to PN argmaxθ i=1 log(Pr(Si∗ | θ)).
In many cases, off-policy learning is vital for practically solving Problem 2.1. Specifically, sampling from our policy directly may generate very few trajectories τ that give a high (or even nonzero) probability Pr(Xi | τ ) for some (or even all) observations Xi . Thus, it makes sense to sample from a distribution that explicitly takes into account X1:N and ensures the sampling of latent states S that have reasonably high probabilities Pr(Xi | S) for various observations Xi . To this end, we use Bayes’ theorem and obtain Pr(τ | Xi , θ) = Pr(Xi | τ ) Pr(τ | θ)/Pr(Xi | θ).
Many approaches exist for solving this problem, such as variational autoencoders (VAE) (Kingma & Welling, 2013), discrete diffusion (Austin et al., 2021), and autoregressive generation (Bengio et al., 2003). Furthermore, if each Xi = Si∗ uniquely determines one trajectory τ , then GReinSS simplifies to autoregressive generation as shown in Theorem A.2
We then have the following theorem. Theorem 3.3. The unbiased variance-minimizing PNoff-policy sampling proposal is q(τ | X1:N , θ) = N1 i=1 Pr(τ | Xi , θ).
Local search baseline: We obtain an additional special case if we know that each indirect observation Xi can be explained by exactly one latent state S ∈ S.
Generally, one cannot directly sample from q(τ | X1:N , θ) but can use heuristics to bias the sampling towards q(τ | X1:N , θ). For example, in the cancer phylogeny application CloMu (Ivanovic & El-Kebir, 2023), the observations Xi directly show which nodes are in the graph generated by
Observation 3.5. If for an observation Xi there is exactly one latent state Ŝi ∈ S such that Pr(Xi | Ŝi ) is non-zero then argmaxS Pr(Xi | S) Pr(S | θ) = argmaxS Pr(Xi | S) = Ŝi for any θ. 4
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
Thus, if for each Xi the probability Pr(Xi | S) is nonzero for only one state S, then Problem 2.2 simplifies to Ŝi = argmaxS Pr(Xi | S). To provide intuition, we implement an extremely simple local search baseline that directly optimizes argmaxS Pr(Xi | S) without using any model Pr(S | θ) with details in Section B.8. This baseline can be effective in cases where Pr(Xi | S) is near zero for all but one state S (as shown in Section 4.1).
sive models as shown in Table 1 (details provided in Sections B.5, B.7 and B.8). Problem 2.2 is solved by sampling S from Pr(S | θ) and then predicting states that maximize Pr(S | θ) Pr(Xi | S) (details in Section B.5). Policy learning baselines: In one special case, GReinSS simplifies to standard policy gradients. Lemma 3.6. Let Pr(Xi | S) = Pr(Xj | S) across all S ∈ S for all i, j ∈ [N ]. Then, GReinSS simplifies to standard PN policy gradients with rewards r′ (τ ) = i=1 Pr(Xi | τ ) and standard reward normalization.
Expectation maximization baselines: Expectation maximization consists of alternating between an E-step of estimating latent states Ŝ1:N and an M-step of optimizing the parameters θ. Specifically, we define q(S | X1:N , θ) = PN 1 i=1 Pr(Xi | S) Pr(S | θ). Intuitively, this distribuN tion is approximately an averaged probability distribution over predicted states Ŝi for i ∈ [N ], since the optimal predicted state Ŝi = argmaxS Pr(Xi | S) Pr(S | θ). As shown in Theorem A.3, expectation maximization alternates between an E-step of estimating latent states S (via the term q(S | X1:N , θ)) and an M-step of maximizing the complete-data log-likelihood expectation value using θ(t+1) = argmaxθ ES∼q(S|X1:N ,θ(t) ) [log(Pr(S | θ))]. The term being maximized only involves S and θ but not X1:N due to the fact that X1:N only depends on θ through S (Theorem A.3). Consequently, the M-step is simply selecting θ to maximize the likelihood of a known distribution of states, similarly to solving Problem 2.1 in the special case where ∗ S1:N is known (as in Observation 3.4).
Intuitively, the denominator Pr(Xi | θ) in the GReinSS reward function (3) helps ensure the trajectory probabilities are balanced across X1:N rather than being overly narrow (as illustrated in Section B.2 and Figure S1). However, when Pr(Xi | θ) is the same for all i, this term becomes irrelevant PN and the reward can simplify to i=1 Pr(Xi | τ ). We define naive policy gradients as a simple ablationPof GReinSS, N where we use the reward function r′ (τ ) = i=1 Pr(Xi | τ ), removing the denominator Pr(Xi | θ). As an alternative baseline, one can better balance probabilities between trajectories without the denominator term Pr(Xi | θ) by using trajectory balance GFlowNets (Malkin et al., 2022). This allows the probability to be balanced across trajectories τ , supporting multiple observations XP i while still using N the predefined reward function r′ (τ ) = i=1 Pr(Xi | τ ). GFlowNets solve Problem 2.1 optimally in the following special case.
For generalized expectation maximization (GEM), the Mstep consists of updating θ to increase the log-likelihood of the distribution of states rather than solving for the optimal θ. Exactly solving GEM cannot be utilized as a baseline method since the E-step cannot be exactly computed for problems with an exponentially large space S of latent states. In principle, the E-step could also be approximated using variational inference by introducing an auxiliary posterior inference model; however, knowing Pr(X | S) enables the simpler approximation of using the current θ to identify latent states Ŝi that maximize Pr(Xi | S) Pr(S | θ) as ∗ approximations of S1:N (i.e. an instance of Problem 2.2). These inferred latent states Ŝ1:N are given equal probability, and other states are given probability 0 rather than marginalizing over the entire exponentially large state space S. Consequently, our approximation of GEM consists of alternating between an approximated E-step of inferring Ŝ1:N and an exact PN M-step of updating θ to increase the log-likelihood i=1 log(Pr(Si∗ | θ)) via gradient descent. Additional details and methodological comparisons with GReinSS are provided in Section B.3.
Lemma 3.7. Let, for each i ∈ [N ], Pr(Xi | τ ) = 1 for exactly one trajectory τ and 0 for all other trajectories. Then, the optimal GFlowNets distribution Pr(τ | θ) is also the optimal solution to Problem 2.1. For this case, Problem 2.2 may be solved by sampling τ from Pr(τ | θ) and then predicting states S = S(τ ) that maximize Pr(S | θ) Pr(Xi | S) (see Section B.5).
4. Results 4.1. Simulations We compare GReinSS to the baseline methods listed in Table 1 on two simulation experiments, where the latent states correspond to distinct combinatorial objects, namely graphs and sets. In the first experiment, our latent states are directed graphs representing some hidden process, and our observations are lists of start and end points of random walks in these directed graphs (Figure 2a). Problem 4.1 (P ROCESS G RAPH I NFERENCE FROM R AN ∗ DOM WALK E NDPOINTS ). Find the directed graphs S1:N ∗ drawn from an unknown distribution Pr (S) given observations X1:N with each Xi composed of k pairs of start and end vertices of k absorbing random walks from Si∗ .
To serve as baselines, this algorithm can easily be combined with standard methods for maximum likelihood generative modeling. We therefore solve GEM using variational autoencoders, discrete diffusion, and autoregres5
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
b
a Si⇤ =
i.e. the harmonic mean of precision and recall of directed edge sets (Section C.2). As shown in Figure 2b and Figure S3, we find that GReinSS consistently achieves the highest median F1 score, followed by the GEM-based baselines using VAE and autoregression. These latter two methods perform very similarly, indicating they may both be effectively solving the GEM M-step, with their performance fundamentally limited by the GEM framework. With the exception of GFlowNets and naive policy gradients, the F1 scores consistently increase with the number k of random walks, showing most approaches can be effective with sufficiently informative observations. Strikingly, for k = 10, GReinSS has the highest median F1 score of 0.891, with all baselines having median F1 scores below 0.55 due to the small amount of information per observation. Despite naive policy gradients and GFlowNets being the most structurally similar baseline methods to GReinSS, they perform very poorly with median F1 scores consistently below 0.55. Specifically, Naive Policy Gradient consistently predicted an empty graph, achieving a high reward for a small number of observations but an F1 score of 0. This demonstrates that while the dynamic reward function is a relatively small algorithmic change to policy gradients, it is absolutely essential for achieving strong results.
-
<latexit sha1_base64="3sR3zkkeUhCxD7Hqy9IbZnVenB0=">AAAB7HicbVBNSwMxEJ31s9avqkcvwSKIh7JbpHosevFY0W0L7VqyabYNTbJLkhXK0t/gxYMiXv1B3vw3pu0etPXBwOO9GWbmhQln2rjut7Oyura+sVnYKm7v7O7tlw4OmzpOFaE+iXms2iHWlDNJfcMMp+1EUSxCTlvh6Gbqt56o0iyWD2ac0EDggWQRI9hYyb9/PO+xXqnsVtwZ0DLxclKGHI1e6avbj0kqqDSEY607XjUxQYaVYYTTSbGbappgMsID2rFUYkF1kM2OnaBTq/RRFCtb0qCZ+nsiw0LrsQhtp8BmqBe9qfif10lNdBVkTCapoZLMF0UpRyZG089RnylKDB9bgoli9lZEhlhhYmw+RRuCt/jyMmlWK16tUru7KNev8zgKcAwncAYeXEIdbqEBPhBg8Ayv8OZI58V5dz7mrStOPnMEf+B8/gBIWY5c</latexit>
𝑋! =
…
k
( , ) ( , )
Figure 2. GReinSS outperforms all baseline methods in simulated process graph inference. a An example latent graph Si∗ and observation Xi consisting of the start and end points of k random walks. b F1 -score as a function of the number k of random walks.
The absorbing random walks are unbiased random walks through the graph Si∗ with a single terminating self-loop on each node, as done in (Wu et al., 2012). This very general problem could represent inferring transportation networks (Biagioni & Eriksson, 2012), biochemical process graphs (Friedman et al., 2000), information transmission graphs (Gomez-Rodriguez et al., 2012), etc. To apply the GReinSS framework, we must recast the above problem by (i) parameterizing a procedure for generating latent state graphs, and (ii) specifying the probability Pr(Xi | S) of generating our indirect observations. Rather than assuming any knowledge on the distribution Pr∗ (S) of graphs, we define θ as the parameters of a neural network that generates graphs S by sequentially adding directed edges to an initially empty graph, ending with a termination action (model details provided in Section B.7).
∗ In the second experiment, our latent states are a family S1:N of N = 100 subsets from a universe U, and observations X1:N are noisy real-valued vectors indicating which elements are more likely to be present in each set (Figure 3a).
Problem 4.2 (S UBSET I NFERENCE FROM N OISY E LE ∗ MENT M EASUREMENTS). Find the family S1:N of subsets ∗ drawn from an unknown distribution Pr (S) on a fixed universe U, given observations X1:N each composed of a vector of length |U| with Xi,j drawn from a known distribution D+ if j ∈ Si∗ and Xi,j drawn from D− otherwise.
The probability Pr(Xi | S) is the product of the start and end probabilities that follow from the inverse shifted Laplacian matrix (L + I)−1 as described in (Wu et al., 2012). As such, we can solve Problem 4.1 naturally using GReinSS by splitting it into a learning problem of identifying parameters θ generating the distribution of graphs Pr(S | θ) (Problem 2.1), and an inference problem of estimating each Si∗ given Xi and θ (Problem 2.2). We generate three simulation instances each composed of N = 1000 observations but a varying number k ∈ {10, 100, 1000} of random walks. To do so, we begin by generating a base graph via Erdős–Rényi with an edge inclusion probability 1/2 (Erdős & Rényi, 1960). Next, we assign a weight to each edge sampled from U (1/4, 1) (with the lower limit set to 1/4 to ensure edge recurrence across latent graphs Si∗ ). We apply weight thresholding with thresholds sampled from U (0, 1) to generate ∗ 1000 latent states S1:1000 . Finally, we generate k random walks recording the pair (v, w) of vertices corresponding to a start node v sampled uniformly at random and the ending node w obtained via an unbiased random walk process (Section C.1).
The above problem where noisy observations indicate which conditions are more likely true for each data point arises in many scientific contexts, including detecting chemicals from noisy mass spectrometry intensities (Aebersold & Mann, 2003), determining active genes from noisy expression data (Schena et al., 1995), etc. This problem fits naturally within the GReinSS framework: Latent states S are generated by adding elements from U iteratively, starting with the empty set ∅, ending with a termination action using a neural network with parameters θ described in Section B.7. The probability distribution Pr(Xi | S) is de+ − rived from Q trivially Q the given− distributions D and D as + j∈S D (Xi,j ) j∈U \S D (Xi,j ). As with Problem 4.1, we solve Problem 4.2 in a two-step fashion by applying GReinSS to first solve Problem 2.1 followed by Problem 2.2. We generate simulation instances with varying noise levels σ ∈ {0.1, 0.2, 0.3, 0.4, 0.5} and varying universe sizes |U| ∈ {10, 100, 1000}. Specifically, we model the distribu-
We measure graph reconstruction accuracy as an F1 score, 6
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
a
b
c
median F1 score of 1.0. However, all baseline methods other than local search and naive policy gradients scale catastrophically to universe sizes of |U| = 1000 (median F1 < 0.4). Meanwhile, for universe sizes of |U| = 1000, GReinSS achieves a median F1 score of 0.938, while local search and naive policy gradients achieve median F1 scores of 0.869 and 0.849, respectively. This demonstrates that simple approaches are sufficient for very small universe sizes, but only GReinSS effectively scales to large vectors. In Figure S2, we analyze policy learning methods without off-policy learning, showing that GReinSS still performs best.
Figure 3. GReinSS outperforms all baseline methods in simulated subset inference. a An example set Si∗ and noisy observation Xi . b GReinSS outperforms baselines for all noise levels σ. c GReinSS scales to large universes U unlike baseline methods.
4.2. RNA Splicing from Short Read RNA-seq Data Cells produce proteins by transcribing DNA into precursor RNA, splicing it into a mature RNA transcript, and translating it into a protein (Alberts et al., 1994). While transcription and translation are mostly deterministic, splicing is a highly variable process that selects and connects contiguous nucleotide segments, called exons. As such, this process yields multiple distinct RNA molecules, called isoforms, from the same gene (Wang et al., 2008) as shown in Figure 4a. Long-read RNA sequencing directly observes full-length isoforms but is expensive (Tilgner et al., 2015; Au et al., 2013). In contrast, inexpensive short-read RNA sequencing produces short (around 100 nucleotide) reads from mature transcripts. As such, isoforms and their proportions must be inferred indirectly from these reads (Mortazavi et al., 2008; Trapnell et al., 2010; Li & Dewey, 2011). Here, we propose to use GReinSS for this problem as follows.
tions D+ and D− as Gaussian distributions N (1, σ 2 ) and N (0, σ 2 ) using the specified noise level σ as the standard deviation. We simulate the latent states Si∗ using a process inspired by dictionary learning (Aharon et al., 2006), in which each set Si∗ is the union of random reusable module subsets (Mairal et al., 2009) with details provided in Section C.3. The state-generating process Pr∗ (S) results in states that exhibit shared structural patterns while maintaining a high degree of random variation. The combination of a widely dispersed and highly randomized state-generation process Pr∗ (S) and very informative observations Xi (especially for small σ) makes off-policy learning extremely valuable. While it is hard to sample exactly from the optimal off-policy sampling distribution Pr(S | X1:N , θ), one can bias the sampling Pr(S | θ) towards Pr(S | X1:N , θ) using the observations X1:N as described in Section C.4. We use off-policy learning for all policy learning methods, including GReinSS, naive policy gradients, and GFlowNets. As shown in Figure S4, the accuracy of GReinSS is not sensitive to the exact off-policy sampling proposal so long as it biases sampling toward Pr(S | X1:N , θ).
Problem 4.3 (I SOFORM P ROPORTION I NFERENCE FROM R EAD C OUNTS). Given a set X1:N of reads aligned to one gene across M sequencing samples, with each read Xi consisting of a sample index and genomic position, estimate the distribution Pr∗ (S) of isoforms in each of the M samples. To apply the GReinSS framework, we first (i) parameterize a procedure for generating latent state isoforms and (ii) specify the probability Pr(Xi | S) of generating our indirect observations. We cast each latent state S as a tuple containing the genomic regions (exons) of the isoform, the sample index of the isoform, and the genomic position of the read Xi produced from this isoform (Figure 4b). The isoform is generated using a neural network with parameters θ by iteratively selecting transitions (junctions) between included contiguous genomic regions (exons), constructing an isoform as described in Section C.5. Our formulation of latent states implies Pr(Xi | S) = 1 if the genomic position and sample index of S and Xi agree, and 0 otherwise.
We assess performance using F1 scores, comparing groundtruth subsets Si∗ to predicted subsets Ŝi . For simulation instances with a universe with |U| = 100 elements, we find that GReinSS consistently outperforms the baseline methods across varying noise levels (Figure 3b and Figure S6). For low σ < 0.3, the best alternatives are naive policy gradients, local search, and GFlowNets, whereas for high σ > 0.3, the best alternatives are VAE and autoregressive. This illustrates how for low σ the key factor is utilizing the observations either via off-policy learning or local search, whereas for high σ the key factor is effectively optimizing Pr(X1:N | θ) using either GReinSS or GEM. Additionally, we analyze the set reconstruction F1 score while varying the size |U| of the universe but keeping σ = 0.3 (Figure 3c and Figure S5). For tiny vectors of |U| = 10 elements, all methods other than naive policy gradients achieve a perfect
To evaluate GReinSS on this problem, we utilize the GTEx database (Lonsdale et al., 2013), containing 17,371 human tissue samples with short-read sequencing data, of which 61 also have matched long-read sequencing data. Specifi7
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
a
exon A
isoform 1: isoform 2:
b exon B
exon C
exon D
exon B exon C
reference gene
c long read (FLAIR):
d
isoform state Si⇤ read observation Xi <latexit sha1_base64="6gV5f0BP07mN1yN7q2Xc3RrGWpc=">AAAB7nicbVDLSgNBEOz1GeMr6tHLYBDEQ8gGiR6DXjxGNA9I1jA7mU2GzM4uM71CWPIRXjwo4tXv8ebfOEn2oIkFDUVVN91dfiyFwXL521lZXVvf2Mxt5bd3dvf2CweHTRMlmvEGi2Sk2z41XArFGyhQ8nasOQ19yVv+6Gbqt564NiJSDziOuRfSgRKBYBSt1Lp/PO+lYtIrFMul8gxkmbgZKUKGeq/w1e1HLAm5QiapMR23EqOXUo2CST7JdxPDY8pGdMA7lioacuOls3Mn5NQqfRJE2pZCMlN/T6Q0NGYc+rYzpDg0i95U/M/rJBhcealQcYJcsfmiIJEEIzL9nfSF5gzl2BLKtLC3EjakmjK0CeVtCO7iy8ukWSm51VL17qJYu87iyMExnMAZuHAJNbiFOjSAwQie4RXenNh5cd6dj3nripPNHMEfOJ8/DgSPaA==</latexit>
isoforms, (ii) optimal transport across isoform proportions, and (iii) a sample average weighted by long-read support (details and an example provided in Section C.6 and Figure S7). Briefly, the isoform prediction error is 0 if correct isoforms and proportions are predicted across all samples and 1 if all predicted isoforms have zero overlap with all long-read-supported isoforms.
<latexit sha1_base64="wGjaJcVIOLgboMBiCoYZvyawOsk=">AAAB6nicbVBNS8NAEJ34WetX1aOXxSJ4KkmR6rHoxWNF+wFtKJvtpF262YTdjVBCf4IXD4p49Rd589+4bXPQ1gcDj/dmmJkXJIJr47rfztr6xubWdmGnuLu3f3BYOjpu6ThVDJssFrHqBFSj4BKbhhuBnUQhjQKB7WB8O/PbT6g0j+WjmSToR3QoecgZNVZ66PR5v1R2K+4cZJV4OSlDjka/9NUbxCyNUBomqNZdr5oYP6PKcCZwWuylGhPKxnSIXUsljVD72fzUKTm3yoCEsbIlDZmrvycyGmk9iQLbGVEz0sveTPzP66YmvPYzLpPUoGSLRWEqiInJ7G8y4AqZERNLKFPc3krYiCrKjE2naEPwll9eJa1qxatVaveX5fpNHkcBTuEMLsCDK6jDHTSgCQyG8Ayv8OYI58V5dz4WrWtOPnMCf+B8/gA14o3F</latexit>
exon D
exon B
reference gene
96% 4% GReinSS:
error: 0.0067
Figure 4c shows the isoform reconstruction for the gene MBD2 in one cultured fibroblast sample (“GTEX-QV440008-SM-447AX”). The ground-truth long-read sequencing detected one isoform with 96% proportion and one isoform with 4% proportion, which we will refer to as the major and minor isoform, respectively. These isoforms are similar with a Jaccard distance of 0.1495. For GReinSS, the major isoform has 92% proportion and the minor isoform has 7% proportion (with other isoforms totaling < 1%). Applying optimal transport gives an error of 0.0067 due to a slight proportion inaccuracy on the two highly similar correct isoforms. For RSEM, the major isoform has 18% proportion, with other highly distinct isoforms totaling 82% proportion, resulting in an error of 0.5374. Taking an average across 61 samples weighted by the long-read sequencing counts gives an error of 0.0069 for GReinSS and 0.471 for RSEM, leading to a difference of 0.0069 − 0.471 = −0.4641, showing that the isoforms identified by GReinSS better match the long-read ground truth. Looking across all genes, Figure 4d and Figure S9 show that GReinSS outperforms RSEM much more frequently than vice versa. Specifically, on 46.6% of genes GReinSS outperforms RSEM by at least 0.05, whereas only on 9.4% of genes RSEM outperforms GReinSS by at least 0.05. This demonstrates how a straightforward application of GReinSS can improve upon the techniques being used in major real-world applications.
92% 7% RSEM:
error: 0.5374 42% 38% 18% gene: MBD2 (chr. 18)
54.16 Mb
54.18 Mb 54.2 Mb 54.22 Mb
Figure 4. GReinSS outperforms RSEM in predicting RNA isoforms from short-read sequencing data. a A diagram of two distinct isoforms of one gene, with solid lines for exons and dashed lines for junctions (transitions between exons). b A diagram of an observation read Xi covering some junction in the isoform Si∗ . c The isoforms detected by long-read sequencing, as well as those predicted by GReinSS and RSEM, are shown for an example gene (MBD2) in one cultured fibroblast sample. Dots show exons and dashed lines show (much longer) junctions. Unlike RSEM, GReinSS reconstructs the same two isoforms identified in the long-read data with very similar proportions. d The GReinSS error minus the RSEM error is shown for all 14,390 genes, with a median difference in errors of −0.0405.
cally, GTEx contains: (i) short-read sequencing junctionoverlapping read counts calculated using the STAR (Dobin et al., 2013) aligner; (ii) isoforms and their proportions estimated from short-read sequencing data using RSEM (Li & Dewey, 2011); and (iii) isoforms and their proportions estimated from long-read sequencing data calculated using FLAIR (Tang et al., 2020). We use (i) as input to GReinSS, (ii) as a baseline method, and (iii) as ground-truth for evaluation. Note that the baseline method RSEM is a commonly used expectation maximization-based algorithm for isoform quantification, which, similarly to GReinSS, uses the same splice-aware read alignments as input. We also include comparisons with the naive policy gradients and GFlowNets ablations in Figure S8, showing that these ablations underperformed both GReinSS and RSEM. We focus our analysis on the M = 61 tissue samples with matched long-read data. Moreover, we restrict our analysis to a total of 14,390 human genes that had a sufficient number (≥ 100) of junction-overlapping reads in the short-read data as well as isoform-covering reads in the long-read data.
5. Conclusion In this work, we introduced GReinSS, a novel framework for learning distributions over discrete latent states from indirect observation data. GReinSS directly optimizes the observation data likelihood without utilizing ground-truth latent states. To achieve this, GReinSS dynamically rescales rewards so that the policy gradient is an unbiased estimator of the observation data log-likelihood gradient. Consequently, GReinSS allows policy learning to be an effective alternative to expectation maximization for combinatorially large latent spaces where an exact implementation of expectation maximization is infeasible. On simulated latent graph and latent set inference problems, GReinSS consistently outperforms baseline methods, including (naive) policy gradients, GFlowNets, local search, and GEM-based implementations of VAEs, autoregressive models, and discrete diffusion. In particular, the poor performance of the naive policy gradients ablation
To evaluate prediction accuracy, we must account for (i) errors in the exons included in each isoform, (ii) errors in predicted proportions, and (iii) the varying long-read sequencing support across samples. We achieve this by utilizing (i) pairwise Jaccard distance between loci covered by 8
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
of GReinSS shows the vital importance of GReinSS’s dynamic reward function. On the biological task of isoform inference from short-read RNA sequencing data, GReinSS outperforms the standard EM-based algorithm RSEM used in the GTEx database. This illustrates the real-world effectiveness of GReinSS on practical problems beyond prior works (Ivanovic & El-Kebir, 2023; 2025).
analysis of biological systems or other forms of experimental data. We do not anticipate that this work introduces novel ethical concerns beyond those broadly applicable to the development of novel machine learning techniques, and we do not identify any specific societal impacts requiring special consideration.
Beyond applying the GReinSS framework to other realworld problems, there is great potential for methodological extensions. First, GReinSS currently uses policy gradients, but could be extended to use Q learning (Watkins & Dayan, 1992) or the actor-critic framework (Sutton et al., 1999). Second, although we derived the optimal off-policy sampling distribution (Theorem 3.3), this and prior work rely on heuristically defined off-policy sampling (Ivanovic & ElKebir, 2023; 2025). Alternatively, one could train an auxiliary network to learn an off-policy sampling distribution that approximates the optimal distribution q(τ | X1:N , θ). Third, an approximate probability function Pr(X | S) of observations X given some state S is currently assumed to be known but could instead contain learned parameters to account for unknown aspects of observation generation. Fourth, more sophisticated forms of multi-environmental modeling could be implemented. The isoform inference application utilizes a fairly simple multi-environment setup, but future extensions could utilize more complex differences in the distributions Pr(S | θ) and Pr(Xi | S) across environments, including the addition of individualized environment-specific parameters. Overall, this work provides a starting point for a wide variety of future applications and extensions.
Code Availability Our implementation of GReinSS is available at https: //github.com/elkebir-group/GReinSS.
References Aebersold, R. and Mann, M. Mass spectrometry-based proteomics. Nature, 422(6928):198–207, 2003. Aharon, M., Elad, M., and Bruckstein, A. K-SVD: An algorithm for designing overcomplete dictionaries for sparse representation. IEEE Transactions on Signal Processing, 54(11):4311–4322, 2006. Alberts, B., Bray, D., Lewis, J., Raff, M., Roberts, K., Watson, J. D., et al. Molecular biology of the cell, volume 3. Garland New York, 1994. Au, K. F., Sebastiano, V., Afshar, P. T., Durruthy, J. D., Lee, L., Williams, B. A., Van Bakel, H., Schadt, E. E., ReijoPera, R. A., Underwood, J. G., et al. Characterization of the human ESC transcriptome by hybrid sequencing. Proceedings of the National Academy of Sciences, 110 (50):E4821–E4830, 2013.
Acknowledgments
Austin, J., Johnson, D. D., Ho, J., Tarlow, D., and Van Den Berg, R. Structured denoising diffusion models in discrete state-spaces. Advances in Neural Information Processing Systems, 34:17981–17993, 2021.
This research was supported by National Science Foundation grant CCF-2046488, the Molecule Maker Lab Institute: An AI Research Institutes program supported by NSF under Award No. 2505932, and the DOE Center for Advanced Bioenergy and Bioproducts Innovation (U.S. Department of Energy, Office of Science, Biological and Environmental Research Program under Award Number DE-SC0018420). Any opinions, findings, and conclusions or recommendations expressed in this publication are those of the author(s) and do not necessarily reflect the views of the U.S. Department of Energy.
Bengio, Y., Ducharme, R., Vincent, P., and Jauvin, C. A neural probabilistic language model. Journal of Machine Learning Research, 3(Feb):1137–1155, 2003. Bengio, Y., Courville, A., and Vincent, P. Representation learning: A review and new perspectives. IEEE Transactions on Pattern Analysis and Machine Intelligence, 35 (8):1798–1828, 2013. Biagioni, J. and Eriksson, J. Inferring road maps from global positioning system traces: Survey and comparative evaluation. Transportation Research Record, 2291(1):61– 71, 2012.
Impact Statement This paper introduces new machine learning methodologies with the primary goal of advancing the field of machine learning. This work is primarily applicable to problems involving latent-variable inference from indirect observations, including scientific and biological data analysis. Potential downstream impacts of this work may include improved
Dempster, A. P., Laird, N. M., and Rubin, D. B. Maximum likelihood from incomplete data via the EM algorithm. Journal of the Royal Statistical Society: Series B (Methodological), 39(1):1–22, 1977. 9
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
Dobin, A., Davis, C. A., Schlesinger, F., Drenkow, J., Zaleski, C., Jha, S., Batut, P., Chaisson, M., and Gingeras, T. R. STAR: ultrafast universal RNA-seq aligner. Bioinformatics, 29(1):15–21, 2013.
Kingma, D. P. and Welling, M. Auto-encoding variational Bayes. arXiv preprint arXiv:1312.6114, 2013. Li, B. and Dewey, C. N. RSEM: accurate transcript quantification from RNA-Seq data with or without a reference genome. BMC Bioinformatics, 12(1):323, 2011.
Erdős, P. and Rényi, A. On the evolution of random graphs. Publications of the Mathematical Institute of the Hungarian Academy of Sciences, 5(1):17–60, 1960.
Lonsdale, J., Thomas, J., Salvatore, M., Phillips, R., Lo, E., Shad, S., Hasz, R., Walters, G., Garcia, F., Young, N., et al. The genotype-tissue expression (GTEx) project. Nature Genetics, 45(6):580–585, 2013.
Felsenstein, J. Evolutionary trees from DNA sequences: a maximum likelihood approach. Journal of Molecular Evolution, 17(6):368–376, 1981.
MacQueen, J. Some methods for classification and analysis of multivariate observations. In Proc of Berkeley Symposium on Mathematical Statistics & Probability, pp. 281–297, 1965.
Friedman, N., Linial, M., Nachman, I., and Pe’er, D. Using Bayesian networks to analyze expression data. In Proceedings of the Fourth Annual International Conference on Computational Molecular Biology, pp. 127–135, 2000.
Mairal, J., Bach, F., Ponce, J., and Sapiro, G. Online dictionary learning for sparse coding. In Proceedings of the 26th Annual International Conference on Machine Learning, pp. 689–696, 2009.
Gomez-Rodriguez, M., Leskovec, J., and Krause, A. Inferring networks of diffusion and influence. ACM Transactions on Knowledge Discovery from Data (TKDD), 5(4): 1–37, 2012.
Malkin, N., Jain, M., Bengio, E., Sun, C., and Bengio, Y. Trajectory balance: Improved credit assignment in GFlowNets. Advances in Neural Information Processing Systems, 35:5955–5967, 2022.
Haarnoja, T., Tang, H., Abbeel, P., and Levine, S. Reinforcement learning with deep energy-based policies. In International Conference on Machine Learning, pp. 1352–1361. PMLR, 2017.
Mnih, A. and Gregor, K. Neural variational inference and learning in belief networks. In International Conference on Machine Learning, pp. 1791–1799. PMLR, 2014.
Hofmann, T. Probabilistic latent semantic indexing. In Proceedings of the 22nd Annual International ACM SIGIR Conference on Research and Development in Information Retrieval, pp. 50–57, 1999.
Mnih, A. and Rezende, D. Variational inference for monte carlo objectives. In International Conference on Machine Learning, pp. 2188–2196. PMLR, 2016.
Ibrahim, S., Mostafa, M., Jnadi, A., Salloum, H., and Osinenko, P. Comprehensive overview of reward engineering and shaping in advancing reinforcement learning applications. IEEE Access, 2024.
Mortazavi, A., Williams, B. A., McCue, K., Schaeffer, L., and Wold, B. Mapping and quantifying mammalian transcriptomes by RNA-seq. Nature Methods, 5(7):621–628, 2008.
Ivanovic, S. and El-Kebir, M. Modeling and predicting cancer clonal evolution with reinforcement learning. Genome Research, 33(7):1078–1088, 2023.
Noé, F., Wu, H., Prinz, J.-H., and Plattner, N. Projected and hidden Markov models for calculating kinetics and metastable states of complex molecules. The Journal of Chemical Physics, 139(18), 2013.
Ivanovic, S. and El-Kebir, M. CNRein: an evolution-aware deep reinforcement learning algorithm for single-cell DNA copy number calling. Genome Biology, 26(1):87, 2025.
Owen, A. B. Monte Carlo theory, methods and examples. https://artowen.su.domains/mc/, 2013. Rabiner, L. R. A tutorial on hidden Markov models and selected applications in speech recognition. Proceedings of the IEEE, 77(2):257–286, 1989.
Ji, W. and Deng, S. Autonomous discovery of unknown reaction pathways from data by chemical reaction neural network. The Journal of Physical Chemistry A, 125(4): 1082–1092, 2021.
Schena, M., Shalon, D., Davis, R. W., and Brown, P. O. Quantitative monitoring of gene expression patterns with a complementary DNA microarray. Science, 270(5235): 467–470, 1995.
Kahn, H. and Marshall, A. W. Methods of reducing sample size in monte carlo computations. Journal of the Operations Research Society of America, 1(5):263–278, 1953.
Sutton, R. S., Barto, A. G., et al. Reinforcement learning: An introduction, volume 1. MIT press Cambridge, 1998. 10
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
Sutton, R. S., McAllester, D., Singh, S., and Mansour, Y. Policy gradient methods for reinforcement learning with function approximation. Advances in Neural Information Processing Systems, 12, 1999.
learning. In International Conference on Machine Learning, pp. 40531–40554. PMLR, 2023.
Tang, A. D., Soulette, C. M., van Baren, M. J., Hart, K., Hrabeta-Robinson, E., Wu, C. J., and Brooks, A. N. Fulllength transcript characterization of SF3B1 mutation in chronic lymphocytic leukemia reveals downregulation of retained introns. Nature Communications, 11(1):1438, 2020. Tilgner, H., Jahanbani, F., Blauwkamp, T., Moshrefi, A., Jaeger, E., Chen, F., Harel, I., Bustamante, C. D., Rasmussen, M., and Snyder, M. P. Comprehensive transcriptome analysis using synthetic long-read sequencing reveals molecular co-association of distant splicing events. Nature Biotechnology, 33(7):736–742, 2015. Trapnell, C., Williams, B. A., Pertea, G., Mortazavi, A., Kwan, G., Van Baren, M. J., Salzberg, S. L., Wold, B. J., and Pachter, L. Transcript assembly and quantification by RNA-seq reveals unannotated transcripts and isoform switching during cell differentiation. Nature Biotechnology, 28(5):511–515, 2010. Wainwright, M. J., Jordan, M. I., et al. Graphical models, exponential families, and variational inference. Foundations and Trends® in Machine Learning, 1(1–2):1–305, 2008. Wang, E. T., Sandberg, R., Luo, S., Khrebtukova, I., Zhang, L., Mayr, C., Kingsmore, S. F., Schroth, G. P., and Burge, C. B. Alternative isoform regulation in human tissue transcriptomes. Nature, 456(7221):470–476, 2008. Watkins, C. J. and Dayan, P. Q-learning. Machine Learning, 8(3):279–292, 1992. Wu, C. J. On the convergence properties of the EM algorithm. The Annals of Statistics, pp. 95–103, 1983. Wu, X.-M., Li, Z., So, A., Wright, J., and Chang, S.-F. Learning with partially absorbing random walks. Advances in Neural Information Processing Systems, 25, 2012. Yan, X., Jeub, L. G., Flammini, A., Radicchi, F., and Fortunato, S. Weight thresholding on complex networks. Physical Review E, 98(4):042304, 2018. Yu, L., Zhang, W., Wang, J., and Yu, Y. Seqgan: Sequence generative adversarial nets with policy gradient. In Proceedings of the AAAI Conference on Artificial Intelligence, volume 31, 2017. Yuan, M., Li, B., Jin, X., and Zeng, W. Automatic intrinsic reward shaping for exploration in deep reinforcement 11
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
A. Proofs (Main Text) Theorem 3.1.
d log(Pr(τ | θ))] with dynamically changing rewards The policy gradient Eτ [r(τ ) dθ
r(τ ) =
N X Pr(Xi | τ ) i=1
.
Pr(Xi | θ)
(4)
d is an unbiased estimator of the gradient dθ log(Pr(X1:N | θ)) of the log-likelihood objective. That is,
d d log(Pr(X1:N | θ)) = Eτ [r(τ ) log(Pr(τ | θ))]. dθ dθ
(5)
Proof. Let T be the set of all trajectories. We then have the following. N N P d X X d X d τ ∈T Pr(Xi | τ ) dθ Pr(τ | θ) P log(Pr(X1 , . . . XN | θ)) = log( Pr(Xi | τ ) Pr(τ | θ)) = dθ dθ i=1 τ ∈T Pr(Xi | τ ) Pr(τ | θ) i=1
(6)
τ ∈T
N XX Pr(Xi | τ ) d
N XX Pr(Xi | τ )
d log(Pr(τ | θ)) Pr(Xi | θ) dθ Pr(Xi | θ) dθ τ ∈T i=1 τ ∈T i=1 ! ! N N X X X Pr(Xi | τ ) d Pr(Xi | τ ) d log(Pr(τ | θ)) = Eτ [ log(Pr(τ | θ))] = Pr(τ | θ) Pr(Xi | θ) dθ Pr(Xi | θ) dθ i=1 i=1 =
Pr(τ | θ) =
Pr(τ | θ)
(7)
(8)
τ ∈T
= Eτ [r(τ )
d log(Pr(τ | θ))] dθ
(Main Text) Theorem 3.3. PN 1 i=1 Pr(τ | Xi , θ). N
(9)
The unbiased variance-minimizing off-policy sampling proposal is q(τ | X1:N , θ) =
Proof. Existing work (Kahn & Marshall, 1953; Owen, 2013) has proven the unbiased variance-minimizing sampling proposal is Pr(τ ) ∝ |r(τ )| Pr(τ | θ). We then have the following. |r(τ )| Pr(τ | θ) =
N X Pr(Xi | τ ) i=1
Pr(Xi | θ)
Pr(τ | θ) =
N X i=1
Pr(τ | Xi , θ) = N q(τ | X1:N , θ).
(10)
Dividing this term by N normalizes it to give a total probability of 1 across all trajectories. Thus, the optimal off-policy sampling proposal is q(τ | X1:N , θ). (Main Text) Lemma 3.4. Let the given observations equal the ground-truth states, i.e., Xi = Si∗ such that Pr(Xi | S) = 1 if S = Xi = Si∗ and 0 otherwise. Then, Problem 2.1 of solving argmaxθ Pr(X1:N | θ) simplifies PN to argmaxθ i=1 log(Pr(Si∗ | θ)). Proof. argmaxθ Pr(X1:N | θ) = argmaxθ log(Pr(X1:N | θ))
(11)
Pr(Xi | θ))
(12)
log(Pr(Xi | θ)))
(13)
log(Pr(Si∗ | θ)))
(14)
= argmaxθ log(
N Y
i=1
= argmaxθ
N X i=1
= argmaxθ
N X i=1
12
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
(Main Text) Lemma 3.6. Let Pr(Xi | S) = Pr(Xj | S) across all S ∈ S for all i, j ∈ [N ]. Then, GReinSS simplifies to PN standard policy gradients with rewards r′ (τ ) = i=1 Pr(Xi | τ ) and standard reward normalization. PN Proof. GReinSS utilizes the reward function r(τ ) = i=1 Pr(Xi | τ )/ Pr(Xi | θ). The naive policy learning reward P N ′ function is r′ (τ ) = i=1 Pr(Xi | τ ). Standard reward normalization modifies this to rN (τ ) = r′ (τ )/Eτ ′ ∼Pr(τ ′ |θ) [r′ (τ ′ )]. Note that for all i, j ∈ [N ] we have Pr(Xi | θ) = Pr(Xj | θ) since Pr(Xi | S) = Pr(Xj | S) for all S ∈ S. Also note that for all i, j ∈ [N ] and all trajectories τ , we have Pr(Xi | τ ) = Pr(Xi | S(τ )) = Pr(Xj | S(τ )) = Pr(Xj | τ ). We then have the following. ′ rN (τ ) =
= = =
r′ (τ ) E[r′ (τ )] PN
(15)
i=1 Pr(Xi | τ ) PN Eτ ′ ∼Pr(τ ′ |θ) [ i=1 Pr(Xi | τ )]
N Pr(X1 | τ ) Eτ ′ ∼Pr(τ ′ |θ) [N Pr(X1 | τ )] Pr(X1 | τ ) Eτ ′ ∼Pr(τ ′ |θ) [Pr(X1 | τ ′ )]
(16) (17) (18)
=
Pr(X1 | τ ) Pr(X1 | θ)
(19)
=
1 X Pr(Xi | τ ) N i=1 Pr(Xi | θ)
(20)
N
=
1 r(τ ). N
(21)
′ Thus, rN (τ ) differs from r(τ ) only by a constant, making them equivalent to optimize.
(Main Text) Lemma 3.7. Let, for each i ∈ [N ], Pr(Xi | τ ) = 1 for exactly one trajectory τ and 0 for all other trajectories. Then, the optimal GFlowNets distribution Pr(τ | θ) is also the optimal solution to Problem 2.1. Proof. For each i ∈ [N ] let τi be the trajectory for which Pr(Xi | τ ) = 1. Let Iτ =τi equal 1 if τ = τi and 0 otherwise. The PN PN GFlowNets reward function is r′ (τ ) = i=1 Pr(Xi | τ ) = i=1 Iτ =τi . The optimal solution for GFlowNets is having the probability Pr(τ | θ) be proportional to the reward r′ (τ ). Consequently, the optimal solution for GFlowNets is achieved by PN setting Pr(τi | θ) = N1 i=1 Iτ =τi . For GReinSS we have the following. argmaxθ Pr(X1 , . . . XN | θ) = argmaxθ = argmaxθ
N X i=1 N X
log(Eτ ∼Pr(τ |θ) [Pr(Xi | τ )])
(22)
log(Eτ ∼Pr(τ |θ) [Iτ =τi ])
(23)
log(Pr(τi | θ)).
(24)
i=1
= argmaxθ
N X i=1
This log-likelihood is maximized by setting Pr(τi | θ) to the empirical distribution in the data Ei∈[N ] Iτ =τi . Thus, the PN optimal solutions for GReinSS and GFlowNets in this special case are both Pr(τi | θ) = N1 i=1 Iτ =τi .
Mini-batching is compatible with GReinSS and still gives an unbiased estimator of the gradient of the observation data log-likelihood as shown below. 13
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
Corollary A.1. Let B ⊂ [N ] be a mini-batch sampled uniformly at random, and define the mini-batch reward rB (τ ) =
X Pr(Xi | τ ) i∈B
Pr(Xi | θ)
.
(25)
Then the policy gradient d N Eτ,B [rB (τ ) log Pr(τ | θ)] |B| dθ
(26)
is an unbiased estimator of the full-data log-likelihood gradient d log Pr(X1:N | θ). dθ
(27)
Proof. d N X Pr(Xi | τ ) d N Eτ,B [rB (τ ) log Pr(τ | θ)] = Eτ,B [ log Pr(τ | θ)] |B| dθ |B| Pr(Xi | θ) dθ
(28)
i∈B
= Eτ [EB [
N X N X Pr(Xi | τ ) d Pr(Xi | τ ) d ] log Pr(τ | θ)] = Eτ [ log Pr(τ | θ)] |B| Pr(Xi | θ) dθ Pr(Xi | θ) dθ i=1
(29)
i∈B
d This is then equal to dθ log Pr(X1:N | θ) by Theorem 3.1.
If each Xi = Si∗ uniquely determines one trajectory τ , then GReinSS simplifies to autoregressive generation as shown below. Lemma A.2. Autoregressive generation using the negative log-likelihood loss is equivalent to GReinSS when for all i ∈ [N ] the probability Pr(Xi | τ ) is non-zero for exactly one known trajectory τi∗ . Proof. For i ∈ [N ], let τi∗ be the one trajectory for which Pr(Xi , τi∗ ) is non-zero. Define δ(τ, τ ′ ) = 1 if τ = τ ′ and 0 otherwise. The off-policy sampling has the following probability. N
Pr(τ | X1 , . . . , XN , θ) =
1 X Pr(τ | Xi , θ) N i=1
(30)
N
=
1 X Pr(τ | θ) Pr(Xi | τ ) N i=1 Pr(Xi | θ)
=
1 X Pr(τ | θ) Pr(Xi | τ ) P N i=1 τ ′ ∈T Pr(Xi | τ ′ ) Pr(τ ′ | θ)
=
1 X Pr(τi∗ | θ) Pr(Xi | τi∗ )δ(τ, τi∗ ) N i=1 Pr(Xi | τi∗ ) Pr(τi∗ | θ)
=
1 X δ(τ, τi∗ ) N i=1
(31)
N
(32)
N
(33)
N
(34)
Thus, optimal importance sampling simply corresponds to uniformly randomly sampling a data point i ∈ [N ], and then selecting the trajectory τi∗ . We note that each τi∗ simply corresponds to a sequence of actions. Therefore, our sampling ∗ procedure corresponds to uniformly randomly selecting sequences from the list τ1∗ , . . . , τN . Additionally, the loss function 14
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
gradient is as follows. N X Y d d log(Pr(X1 , . . . XN | θ)) = log( Pr(Xi | τ ) Pr(τ | θ)) dθ dθ i=1
(35)
τ ∈T
N
d X log(Pr(Xi | τi∗ ) Pr(τi∗ | θ)) dθ i=1 N X d d log(Pr(Xi | τi∗ )) + log(Pr(τi∗ | θ)) = dθ dθ i=1 =
=
N X d
dθ i=1
log(Pr(τi∗ | θ))
(36)
(37)
(38)
This is exactly the negative log-likelihood loss. Uniformly randomly selecting sequences (of actions) and then evaluating the negative log likelihood loss on each sequence is simply autoregressive generation. Theorem A.3. Since X1 , . . . , XN only depend on θ through S, the M-step of expectation maximization simplifies to θ(t+1) = argmaxθ ES∼Pr(S|X1 ,...,XN ,θ(t) ) [log(Pr(S | θ))]. Proof. The standard M-step of expectation maximization is as follows. θ(t+1) = argmaxθ
N X i=1
ESi ∼Pr(S|Xi ,θ(t) ) [log(Pr(Xi , Si | θ))].
(39)
Then, we have the following derivation. θ(t+1) = argmaxθ
N X i=1
= argmaxθ
N X i=1
= argmaxθ
N X i=1
= argmaxθ
N X i=1
EŜi ∼Pr(S|Xi ,θ(t) ) [log(Pr(Xi , Ŝi | θ))].
(40)
EŜi ∼Pr(S|Xi ,θ(t) ) [log(Pr(Xi | Ŝi ) Pr(Ŝi | θ))].
(41)
EŜi ∼Pr(S|Xi ,θ(t) ) [log(Pr(Ŝi | θ))] + ESi ∼Pr(Si |Xi ,θ(t) ) [log(Pr(Xi | Si ))]
(42)
ESi ∼Pr(S|Xi ,θ(t) ) [log(Pr(Si | θ))]
(43)
= argmaxθ N ES∼Pr(S|X1 ,...,XN ,θ(t) ) [log(Pr(S | θ))]
= argmaxθ ES∼Pr(S|X1 ,...,XN ,θ(t) ) [log(Pr(S | θ))].
(44) (45)
B. Methodological and Implementation Details B.1. Multiple Pools of Observations Simple pooled observations: If multiple pieces of data are generated from each state S, one can group together these sub-observations into one full observation Xi for each state. For instance, if each Si∗ is a transportation network, each Xi may consist of a list of observed transportation paths (sub-observations) through that network. As the number of sub-observations per observation Xi increases, each observation Xi provides more information about the latent state Si∗ , and the task of estimating the latent state Si∗ becomes easier (illustrated in Section 4.1), and the probability distribution Pr(Xi | S) may approach zero for many incorrect states S ̸= Si∗ . For the P ROCESS G RAPH I NFERENCE problem in Section 4.1, each state is a directed graph, and observations consist of lists of start and end points of random walks through 15
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
a
Pr(X | S) <latexit sha1_base64="sJrlw/TCB9RXSdJ9dDQYNhuq8wY=">AAAB9HicbVBNT8JAEJ3iF+IX6tHLRmKCF9ISgx6JXjxiFCShDdlut7Bht627WxLS8Du8eNAYr/4Yb/4bF+hBwZdM8vLeTGbm+QlnStv2t1VYW9/Y3Cpul3Z29/YPyodHHRWnktA2iXksuz5WlLOItjXTnHYTSbHwOX30Rzcz/3FMpWJx9KAnCfUEHkQsZARrI3luS1a7yBUsQPfn/XLFrtlzoFXi5KQCOVr98pcbxCQVNNKEY6V6Tj3RXoalZoTTaclNFU0wGeEB7RkaYUGVl82PnqIzowQojKWpSKO5+nsiw0KpifBNp8B6qJa9mfif10t1eOVlLEpSTSOyWBSmHOkYzRJAAZOUaD4xBBPJzK2IDLHERJucSiYEZ/nlVdKp15xGrXF3UWle53EU4QROoQoOXEITbqEFbSDwBM/wCm/W2Hqx3q2PRWvBymeO4Q+szx8Cv5D5</latexit>
c
b
S1 S2 S3 <latexit sha1_base64="s61jL5JraE3myXbdkB8HCadn+s8=">AAAB6nicbVDLTgJBEOzFF+IL9ehlIjHxRFhi0CPRi0cM8khgQ2aHWZgwO7uZ6TUhGz7BiweN8eoXefNvHGAPClbSSaWqO91dfiyFwUrl28ltbG5t7+R3C3v7B4dHxeOTtokSzXiLRTLSXZ8aLoXiLRQoeTfWnIa+5B1/cjf3O09cGxGpR5zG3AvpSIlAMIpWajYH1UGxVClXFiDrxM1ICTI0BsWv/jBiScgVMkmN6bnVGL2UahRM8lmhnxgeUzahI96zVNGQGy9dnDojF1YZkiDSthSShfp7IqWhMdPQt50hxbFZ9ebif14vweDGS4WKE+SKLRcFiSQYkfnfZCg0ZyinllCmhb2VsDHVlKFNp2BDcFdfXiftatmtlWsPV6X6bRZHHs7gHC7BhWuowz00oAUMRvAMr/DmSOfFeXc+lq05J5s5hT9wPn8A2tmNiQ==</latexit>
X1
0.5
0
0
S(⌧2 ) = S2
X2
0
0.3 0.2
S(⌧3 ) = S3
<latexit sha1_base64="eE3EO2fMwKDgWSP9kaNJHS2kG/E=">AAAB+HicbVDLSsNAFJ34rPXRqEs3g0Wom5K0Ut0IRTcuK7UPaEOYTCft0MkkzEOooV/ixoUibv0Ud/6N0zYLbT1w4XDOvdx7T5AwKpXjfFtr6xubW9u5nfzu3v5BwT48astYC0xaOGax6AZIEkY5aSmqGOkmgqAoYKQTjG9nfueRCElj/qAmCfEiNOQ0pBgpI/l2oVnqK6T96jm8hk2/6ttFp+zMAVeJm5EiyNDw7a/+IMY6IlxhhqTsuZVEeSkSimJGpvm+liRBeIyGpGcoRxGRXjo/fArPjDKAYSxMcQXn6u+JFEVSTqLAdEZIjeSyNxP/83pahVdeSnmiFeF4sSjUDKoYzlKAAyoIVmxiCMKCmlshHiGBsDJZ5U0I7vLLq6RdKbu1cu3+oli/yeLIgRNwCkrABZegDu5AA7QABho8g1fwZj1ZL9a79bFoXbOymWPwB9bnD0HekYw=</latexit>
r(⌧1 ) Pr(⌧1 | ✓) = 0.5 Pr(⌧2 | ✓) = 0.5 r(⌧2 ) Pr(⌧3 | ✓) = 0 <latexit sha1_base64="klLvurhOxuqf4jC9/wsBU2j1y3U=">AAACBnicbVDLSsNAFJ34rPUVdSnCYBHqpiRFqxuh6MZlBfuAJoTJZNoOnUzCzI1QSldu/BU3LhRx6ze482+cPhbaeuDC4Zx7ufeeMBVcg+N8W0vLK6tr67mN/ObW9s6uvbff0EmmKKvTRCSqFRLNBJesDhwEa6WKkTgUrBn2b8Z+84EpzRN5D4OU+THpSt7hlICRAvvIq6miByQLXOzFPMIe9BiQU3yFndJ5YBeckjMBXiTujBTQDLXA/vKihGYxk0AF0brtllPwh0QBp4KN8l6mWUpon3RZ21BJYqb94eSNET4xSoQ7iTIlAU/U3xNDEms9iEPTGRPo6XlvLP7ntTPoXPpDLtMMmKTTRZ1MYEjwOBMcccUoiIEhhCpubsW0RxShYJLLmxDc+ZcXSaNcciulyt1ZoXo9iyOHDtExKiIXXaAqukU1VEcUPaJn9IrerCfrxXq3PqatS9Zs5gD9gfX5A+VRltc=</latexit>
<latexit sha1_base64="LYGySeQJ874qa070Z8ZKBxxJRO4=">AAAB8HicbVA9TwJBEJ3DL8Qv1NLmIjHBhtwRg5ZEG0tMBDFwIXvLHmzY3bvszpkQwq+wsdAYW3+Onf/GBa5Q8CWTvLw3k5l5YSK4Qc/7dnJr6xubW/ntws7u3v5B8fCoZeJUU9aksYh1OySGCa5YEzkK1k40IzIU7CEc3cz8hyemDY/VPY4TFkgyUDzilKCVHnW5iyTt+ee9YsmreHO4q8TPSAkyNHrFr24/pqlkCqkgxnT8aoLBhGjkVLBpoZsalhA6IgPWsVQRyUwwmR88dc+s0nejWNtS6M7V3xMTIo0Zy9B2SoJDs+zNxP+8TorRVTDhKkmRKbpYFKXCxdidfe/2uWYUxdgSQjW3t7p0SDShaDMq2BD85ZdXSata8WuV2t1FqX6dxZGHEziFMvhwCXW4hQY0gYKEZ3iFN0c7L86787FozTnZzDH8gfP5A+xYj9o=</latexit>
Pr(⌧1 | ✓) => <latexit sha1_base64="vmnijwdNZNpUm6dVbRkKtgoT2UQ=">AAACAHicbVDLSsNAFJ3UV62vqAsXbgaLUDclKVJdFt24rGAf0IQwmUzaoZMHMzdCCd34K25cKOLWz3Dn3zhts9DWAxcO59zLvff4qeAKLOvbKK2tb2xulbcrO7t7+wfm4VFXJZmkrEMTkci+TxQTPGYd4CBYP5WMRL5gPX98O/N7j0wqnsQPMEmZG5FhzENOCWjJM0+ctqw5QDLPxk7EA+zAiAG58MyqVbfmwKvELkgVFWh75pcTJDSLWAxUEKUGdiMFNycSOBVsWnEyxVJCx2TIBprGJGLKzecPTPG5VgIcJlJXDHiu/p7ISaTUJPJ1Z0RgpJa9mfifN8ggvHZzHqcZsJguFoWZwJDgWRo44JJREBNNCJVc34rpiEhCQWdW0SHYyy+vkm6jbjfrzfvLauumiKOMTtEZqiEbXaEWukNt1EEUTdEzekVvxpPxYrwbH4vWklHMHKM/MD5/ADerlYs=</latexit>
<latexit sha1_base64="Sz8W6X/077J61ATFCuzHb30GcPA=">AAAB+HicbVDLSsNAFJ34rPXRqEs3g0Wom5IEqW6EohuXldoHtCFMppN26GQS5iHU0C9x40IRt36KO//GaZuFth64cDjnXu69J0wZlcpxvq219Y3Nre3CTnF3b/+gZB8etWWiBSYtnLBEdEMkCaOctBRVjHRTQVAcMtIJx7czv/NIhKQJf1CTlPgxGnIaUYyUkQK71Kz0FdKBdw6vYTPwArvsVJ054Cpxc1IGORqB/dUfJFjHhCvMkJQ910uVnyGhKGZkWuxrSVKEx2hIeoZyFBPpZ/PDp/DMKAMYJcIUV3Cu/p7IUCzlJA5NZ4zUSC57M/E/r6dVdOVnlKdaEY4XiyLNoErgLAU4oIJgxSaGICyouRXiERIIK5NV0YTgLr+8Stpe1a1Va/cX5fpNHkcBnIBTUAEuuAR1cAcaoAUw0OAZvII368l6sd6tj0XrmpXPHIM/sD5/AD7PkYo=</latexit>
<latexit sha1_base64="k/0rEMpBTAWCKauBTnv6DcrpCSw=">AAAB6nicbVBNS8NAEJ34WetX1aOXxSJ4KkmR6rHoxWNF+wFtKJvtpF262YTdjVBCf4IXD4p49Rd589+4bXPQ1gcDj/dmmJkXJIJr47rfztr6xubWdmGnuLu3f3BYOjpu6ThVDJssFrHqBFSj4BKbhhuBnUQhjQKB7WB8O/PbT6g0j+WjmSToR3QoecgZNVZ66PS9fqnsVtw5yCrxclKGHI1+6as3iFkaoTRMUK27XjUxfkaV4UzgtNhLNSaUjekQu5ZKGqH2s/mpU3JulQEJY2VLGjJXf09kNNJ6EgW2M6JmpJe9mfif101NeO1nXCapQckWi8JUEBOT2d9kwBUyIyaWUKa4vZWwEVWUGZtO0YbgLb+8SlrViler1O4vy/WbPI4CnMIZXIAHV1CHO2hAExgM4Rle4c0Rzovz7nwsWtecfOYE/sD5/AHg842N</latexit>
<latexit sha1_base64="EC7/gYwpgX/oIDhha3wMSdBRsBg=">AAAB6nicbVBNS8NAEJ34WetX1aOXxSJ4KkmR6rHoxWNF+wFtKJvtpl262YTdiVBCf4IXD4p49Rd589+4bXPQ1gcDj/dmmJkXJFIYdN1vZ219Y3Nru7BT3N3bPzgsHR23TJxqxpsslrHuBNRwKRRvokDJO4nmNAokbwfj25nffuLaiFg94iThfkSHSoSCUbTSQ6df7ZfKbsWdg6wSLydlyNHol756g5ilEVfIJDWm61UT9DOqUTDJp8VeanhC2ZgOeddSRSNu/Gx+6pScW2VAwljbUkjm6u+JjEbGTKLAdkYUR2bZm4n/ed0Uw2s/EypJkSu2WBSmkmBMZn+TgdCcoZxYQpkW9lbCRlRThjadog3BW355lbSqFa9Wqd1flus3eRwFOIUzuAAPrqAOd9CAJjAYwjO8wpsjnRfn3flYtK45+cwJ/IHz+QPid42O</latexit>
d
<latexit sha1_base64="dQPKHm+UsYmUQWwMPjVsIcjxEkg=">AAAB+HicbVDLSgNBEJz1GeMjqx69DAYhXkI2SPQiBL14jMQ8IFmW2clsMmT2wUyPEJd8iRcPinj1U7z5N06SPWhiQUNR1U13l58IrqBS+bbW1jc2t7ZzO/ndvf2Dgn141FaxlpS1aCxi2fWJYoJHrAUcBOsmkpHQF6zjj29nfueRScXj6AEmCXNDMox4wCkBI3l2oVnqA9Gec46vcdNzPLtYKVfmwKvEyUgRZWh49ld/EFMdsgioIEr1nGoCbkokcCrYNN/XiiWEjsmQ9QyNSMiUm84Pn+IzowxwEEtTEeC5+nsiJaFSk9A3nSGBkVr2ZuJ/Xk9DcOWmPEo0sIguFgVaYIjxLAU84JJREBNDCJXc3IrpiEhCwWSVNyE4yy+vkna17NTKtfuLYv0miyOHTtApKiEHXaI6ukMN1EIUafSMXtGb9WS9WO/Wx6J1zcpmjtEfWJ8/O8CRiA==</latexit>
S(⌧1 ) = S1
<latexit sha1_base64="zbu/3tFI6X4m3XZ3S/TzCuDrrM8=">AAAB6nicbVDLTgJBEOz1ifhCPXqZSEw8kV006JHoxSMGeSSwIbNDL0yYnd3MzJoQwid48aAxXv0ib/6NA+xBwUo6qVR1p7srSATXxnW/nbX1jc2t7dxOfndv/+CwcHTc1HGqGDZYLGLVDqhGwSU2DDcC24lCGgUCW8Hobua3nlBpHstHM07Qj+hA8pAzaqxUr/cue4WiW3LnIKvEy0gRMtR6ha9uP2ZphNIwQbXueOXE+BOqDGcCp/luqjGhbEQH2LFU0gi1P5mfOiXnVumTMFa2pCFz9ffEhEZaj6PAdkbUDPWyNxP/8zqpCW/8CZdJalCyxaIwFcTEZPY36XOFzIixJZQpbm8lbEgVZcamk7cheMsvr5JmueRVSpWHq2L1NosjB6dwBhfgwTVU4R5q0AAGA3iGV3hzhPPivDsfi9Y1J5s5gT9wPn8A3F2Nig==</latexit>
<latexit sha1_base64="GgfTZ+cuS+yTCQ8FTbc88lF50P4=">AAAB6nicbVDLTgJBEOzFF+IL9ehlIjHxRFhi0CPRi0cM8khgQ2aHWZgwO7uZ6TUhGz7BiweN8eoXefNvHGAPClbSSaWqO91dfiyFwUrl28ltbG5t7+R3C3v7B4dHxeOTtokSzXiLRTLSXZ8aLoXiLRQoeTfWnIa+5B1/cjf3O09cGxGpR5zG3AvpSIlAMIpWajYH7qBYqpQrC5B14makBBkag+JXfxixJOQKmaTG9NxqjF5KNQom+azQTwyPKZvQEe9ZqmjIjZcuTp2RC6sMSRBpWwrJQv09kdLQmGno286Q4tisenPxP6+XYHDjpULFCXLFlouCRBKMyPxvMhSaM5RTSyjTwt5K2JhqytCmU7AhuKsvr5N2tezWyrWHq1L9NosjD2dwDpfgwjXU4R4a0AIGI3iGV3hzpPPivDsfy9ack82cwh84nz/ZVY2I</latexit>
<latexit sha1_base64="qka/9JwNeQOLGjeduhQwVlwS9n8=">AAACBnicbVDLSsNAFJ34rPUVdSnCYBHqpiRFqxuh6MZlBfuAJoTJZNoOnUzCzI1QSldu/BU3LhRx6ze482+cPhbaeuDC4Zx7ufeeMBVcg+N8W0vLK6tr67mN/ObW9s6uvbff0EmmKKvTRCSqFRLNBJesDhwEa6WKkTgUrBn2b8Z+84EpzRN5D4OU+THpSt7hlICRAvvIq6miByQLytiLeYQ96DEgp/gKO6XzwC44JWcCvEjcGSmgGWqB/eVFCc1iJoEKonXbLafgD4kCTgUb5b1Ms5TQPumytqGSxEz7w8kbI3xilAh3EmVKAp6ovyeGJNZ6EIemMybQ0/PeWPzPa2fQufSHXKYZMEmnizqZwJDgcSY44opREANDCFXc3IppjyhCwSSXNyG48y8vkka55FZKlbuzQvV6FkcOHaJjVEQuukBVdItqqI4oekTP6BW9WU/Wi/VufUxbl6zZzAH6A+vzB+boltg=</latexit>
<latexit sha1_base64="EuIIhsCd8DP5zp9bDLLVHVAtpsw=">AAAB8HicbVA9TwJBEJ3DL8Qv1NJmIzHBhnDEoCXRxhITQQxcyN6yBxt27y67cybkwq+wsdAYW3+Onf/GBa5Q8CWTvLw3k5l5fiyFwWr128mtrW9sbuW3Czu7e/sHxcOjtokSzXiLRTLSHZ8aLkXIWyhQ8k6sOVW+5A/++GbmPzxxbUQU3uMk5p6iw1AEglG00qMu95Am/dp5v1iqVqpzkFXiZqQEGZr94ldvELFE8RCZpMZ03VqMXko1Cib5tNBLDI8pG9Mh71oaUsWNl84PnpIzqwxIEGlbIZK5+nsipcqYifJtp6I4MsveTPzP6yYYXHmpCOMEecgWi4JEEozI7HsyEJozlBNLKNPC3krYiGrK0GZUsCG4yy+vknat4tYr9buLUuM6iyMPJ3AKZXDhEhpwC01oAQMFz/AKb452Xpx352PRmnOymWP4A+fzB+3dj9s=</latexit>
Pr(⌧2 | ✓) => <latexit sha1_base64="pxl/qMKfWN45bl4c9C/NnYUFM9A=">AAACAHicbVDLSsNAFJ3UV62vqAsXbgaLUDclKVJdFt24rGAf0IQwmUzaoZMHMzdCCd34K25cKOLWz3Dn3zhts9DWAxcO59zLvff4qeAKLOvbKK2tb2xulbcrO7t7+wfm4VFXJZmkrEMTkci+TxQTPGYd4CBYP5WMRL5gPX98O/N7j0wqnsQPMEmZG5FhzENOCWjJM0+ctqw5QDKvgZ2IB9iBEQNy4ZlVq27NgVeJXZAqKtD2zC8nSGgWsRioIEoN7EYKbk4kcCrYtOJkiqWEjsmQDTSNScSUm88fmOJzrQQ4TKSuGPBc/T2Rk0ipSeTrzojASC17M/E/b5BBeO3mPE4zYDFdLAozgSHBszRwwCWjICaaECq5vhXTEZGEgs6sokOwl19eJd1G3W7Wm/eX1dZNEUcZnaIzVEM2ukItdIfaqIMomqJn9IrejCfjxXg3PhatJaOYOUZ/YHz+ADk8lYw=</latexit>
<latexit sha1_base64="3qKHHs6wbJ77MmChf9YlKR14hGg=">AAACBHicbVDLSsNAFJ34rPUVddnNYBHqpiRVqhuh6MZlBfuAJoTJZNoOnTyYuRFK6MKNv+LGhSJu/Qh3/o3TNgttPXDhcM693HuPnwiuwLK+jZXVtfWNzcJWcXtnd2/fPDhsqziVlLVoLGLZ9YligkesBRwE6yaSkdAXrOOPbqZ+54FJxePoHsYJc0MyiHifUwJa8syS05QVB0jqnWEn5AF2YMiAnOIrbHlm2apaM+BlYuekjHI0PfPLCWKahiwCKohSPbuWgJsRCZwKNik6qWIJoSMyYD1NIxIy5WazJyb4RCsB7sdSVwR4pv6eyEio1Dj0dWdIYKgWvan4n9dLoX/pZjxKUmARnS/qpwJDjKeJ4IBLRkGMNSFUcn0rpkMiCQWdW1GHYC++vEzatapdr9bvzsuN6zyOAiqhY1RBNrpADXSLmqiFKHpEz+gVvRlPxovxbnzMW1eMfOYI/YHx+QP2jpZi</latexit>
Figure S1. A simple intuitive example applying GReinSS. a The problem setup is defined by the set of states S1 , S2 , S3 , the set of observations X1 , X2 , and the probability function Pr(X | S) defined on these states and observations. b Policy learning is applied by having the states S1 , S2 , S3 generated by trajectories τ1 , τ2 , τ3 . c Increasing the probability Pr(τ | θ) of either trajectory τ1 or τ2 , results in a decreased dynamic reward r(τ ) for that trajectory. d The dynamic rewards result in an optimal solution where the probabilities Pr(τ1 | θ) and Pr(τ2 | θ) are balanced, maximizing the probability Pr(X1:N | θ) = Pr(X1 | θ) Pr(X2 | θ). However, Pr(τ3 | θ) = 0 since the trajectory τ3 is not needed to maximize Pr(X1:N | θ).
the graph. The number of sub-observations is varied, showing how this impacts the accuracy of each method. In simulations, modifying np allows us to test how different methods perform when the amount of information about each ground-truth state Si∗ in each observation Xi varies. Multiple environments: In some cases, one may have several environments v, each with their own distribution of ∗ ∗ latent states Pr∗v (S), list of ground-truth latent states Sv,1 , . . . Sv,N v , distribution of observations Prv (X | S), and list v v of observations X1 , . . . XN v . This may appear to be a generalization, but in fact, it is a special case of the original Problems 2.1 and 2.2. We simply append the environment number v to each observation and state. Specifically, we define Pr((X, v) | (S, v ′ )) = Prv (X | S) for v = v ′ , and 0 otherwise. Similarly, define Pr∗ ((S, v)) = Pr∗v (S), and define Pr((S, v) | θ) as the probability of the latent state S in environment v given model parameters θ. We then solve Problems 2.1 M and 2.2 on the full set of observations (X11 , 1), (X21 , 1), . . . (XN , nenv ) where nenv is the number of environments. This technique is used for modeling RNA splicing in Section 4.2, with each sample treated as an environment. B.2. Intuitive Example Demonstrating Adaptive Rewards Imagine a simple toy problem with three possible states S1 , S2 , S3 and two observations X1 , X2 , with Pr(X1 | S1 ) = 0.5, Pr(X2 | S2 ) = 0.3, Pr(X2 | S3 ) = 0.2, and Pr(Xi | Sj ) = 0 for all other i ∈ [2], j ∈ [3] (Figure S1a). For simplicity, imagine each Si has its own unique trajectory τi (Figure S1b). The fixed reward function without GReinSS’s rescaling Pstate N is r′ (τ ) = i=1 Pr(Xi | τ ). Thus, r′ (τ1 ) = 0.5, r′ (τ2 ) = 0.3, r′ (τ3 ) = 0.2. Since τ1 achieves the highest reward, standard policy gradients would solve for the optimal solution where Pr(τ1 | θ) = 1, Pr(τ2 | θ) = Pr(τ3 | θ) = 0. Thus, Pr(X2 | θ) = 0, and Pr(X1 , X2 | θ) = 0. For GReinSS we have the following rewards. 1 Pr(X1 | τ1 ) = Pr(X1 | τ1 ) Pr(τ1 | θ) Pr(τ1 | θ)
(46)
r(τ2 ) =
Pr(X2 | τ2 ) Pr(X2 | τ2 ) Pr(τ2 | θ) + Pr(X2 | τ3 ) Pr(τ3 | θ)
(47)
r(τ3 ) =
Pr(X2 | τ3 ) Pr(X2 | τ2 ) Pr(τ2 | θ) + Pr(X2 | τ3 ) Pr(τ3 | θ)
(48)
r(τ1 ) =
We have r(τ2 ) > r(τ3 ) for any θ which results in Pr(τ3 | θ) = 0. Thus, we have the following simplification. r(τ2 ) =
Pr(X2 | τ2 ) 1 = Pr(X2 | τ2 ) Pr(τ2 | θ) Pr(τ2 | θ)
(49)
When Pr(τ1 | θ) > Pr(τ2 | θ) then r(τ2 ) > r(τ1 ) and when Pr(τ1 | θ) < Pr(τ2 | θ) then r(τ2 ) < r(τ1 ) (Figure S1c). The equilibrium solution is r(τ1 ) = r(τ2 ) when Pr(τ1 | θ) = Pr(τ2 | θ). Thus, GReinSS solves for Pr(τ1 | θ) = Pr(τ2 | θ) = 0.5, and Pr(τ3 | θ) = 0 (Figure S1d). Thus, Pr(X1 | θ) = 0.52 = 0.25, and Pr(X2 | θ) = 0.5 · 0.3 = 0.15. Consequently, Pr(X1 , X2 | θ) = 0.25 · 0.15 = 0.0375. 16
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
B.3. Methodological comparison of GReinSS with expectation maximization GReinSS and generalized expectation maximization (GEM) both have the goal of maximizing the probability Pr(X1 , . . . , XN | θ). GEM consists of alternating between an E-step of estimating the latent states and an M-step of optimizing θ to maximize an expectation value given the estimated latent states. Specifically, for Problem 2.1, GEM cannot be solved exactly, so the E-step is approximated by predicting latent states Ŝ1 , . . . , ŜN with Ŝi = argmaxS Pr(S | θ) Pr(Xi | S). PN Then, since Xi only depends on θ through S, the M-step simplifies to optimizing θ to maximize the i=1 log(Pr(Ŝi | θ)) as shown in Theorem A.3. With this approximation, the first step of GReinSS and GEM both consist of sampling S from Pr(S | θ). For GReinSS, the next step is using Pr(Xi | S) to calculate rewards for the sampled states and modifying θ using policy gradients based on these rewards. For GEM, the next step is using Pr(Xi | S) to estimate Ŝ1 , . . . ŜN using the sampled states and then PN modifying θ to maximize the probability i=1 log(Pr(Ŝi | θ)). Both GReinSS and GEM repeatedly alternate between these PN steps until convergence of i=1 log(Pr(Ŝi | θ)). Since the M-step of GEM consists of maximum likelihood generative modeling, it can utilize any machine learning method designed for this task, including autoregressive models, variational autoencoders, and discrete diffusion. B.4. Prior Applications of GReinSS Two previous papers have utilized special cases of the GReinSS technique without formulating the full generalized GReinSS procedure (Ivanovic & El-Kebir, 2023; 2025). One of these papers introduces CLoMu, a method for modeling and predicting cancer clonal evolution (Ivanovic & El-Kebir, 2023). CloMu specifically defines the latent states as unobserved phylogenetic (evolutionary) trees representing the clonal evolution of SNV mutations in the tumor of a specific patient. These phylogenetic trees are generated by starting with an unmutated cell, and iteratively adding SNV mutations to populations of cells in order to form a phylogeny tree of populations of cells with different mutations. Thus, the actions generating the trajectory correspond to adding mutations to populations of cells (clones) in order to form new populations of cells (clones) with new mutations. The combinatorial structure of phylogeny trees makes their generation via policy learning very natural. Phylogeny trees are never directly observed. Instead, DNA sequencing enables noisy measurements of mutations on existing populations of cells, from which sets of possible phylogeny trees can be derived. The observations used consist of possible sets of phylogeny trees for each patient, and are generated from the DNA sequencing data. The second prior application of GReinSS is CNRein (Ivanovic & El-Kebir, 2025), a method for inferring CNV mutations on individual cells from single-cell DNA sequencing of tumors. A latent state consists of the set of CNV mutations present in an individual cell in the form of a copy number profile that indicates the total number of copies of each region of the genome, as well as the haplotype that the copies correspond to. Mathematically, each latent state is represented by a sequence of pairs of integers in which each integer in the pair represents one of the two haplotypes. These latent states are generated by iteratively adding CNV mutations to a cell, starting with a normal cell. Thus, actions correspond to selecting a CNV mutation to add, represented in terms of the genomic region covered by the CNV, the haplotype of the CNV, and the copy number change of the CNV. The observation data consist of a processed form of the DNA sequencing data of each cell. One component of this sequencing data is the read depth, which is defined as the number of DNA sequencing reads that are determined to originate from each genomic region. Another component is the B-allele frequency, which is defined as the proportion of DNA sequencing reads that are determined to originate from each haplotype for some given genomic region. Mathematically, each of these is represented by a list of real values with one real value for each genomic position. Although these measurements are biologically complex, the relationship between these observations and the latent states is modeled with a simple Gaussian distribution. For instance, the read depth of a genomic position (a component of Xi ) is modeled by a Gaussian distribution with a mean equal to the integer total copy number of that genomic position (a component of S). This formulation is very similar to our latent set reconstruction problem, with the modification of allowing for integer-valued vectors rather than binary vector representations. B.5. Latent State Inference Given some model parameters θ, and the ability to sample states from the distribution Pr(S | θ), we describe the procedure ∗ for solving Problem 2.2 and predicting states Sˆ1 , . . . , SˆN with the goal of matching the ground-truth latent states S1∗ , . . . , SN . ′ ′ Let IS=S ′ be defined as 1 if S = S and 0 otherwise. Note that Pr(Xi | S) Pr(S | θ) = ES∼Pr(S|θ) [IS=S ′ Pr(Xi | S )] can be estimated via sampling from Pr(Xi | S) for all methods regardless of whether Pr(Xi | S) has a closed-form expression. Thus, for all methods we sample s from Pr(S | θ) and predict the state Ŝi = argmaxS ′ ES∼Pr(S|θ) [IS=S ′ Pr(Xi | S ′ )]. 17
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
For methods with off-policy sampling, we apply the importance sampling correction to this expectation value. During the training of GEM-based methods, we sample a batch size of 1000 latent states to calculate the E-step estimates of Ŝ1 , . . . , ŜN prior to the M-step update of model parameters θ. For all methods, during the final prediction of latent states, we sample 100,000 latent states from Pr(S | θ) to calculate Ŝ1 , . . . , ŜN . B.6. Importance Sampling After modifying the sampling procedure, one simply has to use importance sampling to account for this in policy gradient. Let π ′ be the modified sampling procedure that utilizes the observations X1 , . . . , XN to generate trajectories τ with high probabilities Pr(Xi | τ ). Then, applying importance sampling gives the below equation d Pr(τ | θ) d log(Pr(X1 , . . . XN | θ)) = Eτ ∼π′ [r(τ ) log(Pr(τ | θ))]. dθ Pr(τ | π ′ ) dθ
(50)
Thus, off-policy sampling of τ gives the correct policy gradient if one includes the importance sampling correction Pr(τ | θ)/ Pr(τ | π ′ ). B.7. Model Architectures All models utilize relatively similar neural network architectures, so that the differences in results between models result from their different training procedure rather than their model architecture. All policy-learning-based methods, including GReinSS, naive policy gradients, and trajectory balance GFlowNets, use an identical model architecture. This model consists of a two-layer fully connected neural network with 50 hidden neurons and the PyTorch leakyRelu non-linearity. The autoregressive model uses a similar architecture consisting of a two-layer neural network with autoregressive masking and consequently uses a number of hidden neurons equal to the number of tokens in the input. The variational autoencoder uses two two-layer neural networks with 50 hidden neurons for both the encoder and decoder. The size of the latent representation inside the auto-encoder is also 50. The discrete diffusion model also uses a two-layer neural network with 50 hidden neurons. For the discrete diffusion model, 20 timesteps are used. Timestep values are concatenated with the input after first applying a cosine-based transformation. Specifically, the timestep t is embedded into a 20 dimensional vector where the ith element is defined as cos(t · i · π/20). The local search algorithm does not use a trained model. Latent states are instead represented as a binary vector, and for each Xi each iteration changes one element in the binary vector in order to maximize Pr(Xi | s). B.8. Additional Baseline Details The autoregressive model is trained using the standard negative log-likelihood loss. The variational autoencoder is trained using the standard evidence lower bound (ELBO) loss. The discrete diffusion model is trained using the standard denoising diffusion loss. Trajectory balance GFlowNets use the standard trajectory balance loss. Naive policy gradients uses standard policy gradients with no augmentation of the reward function. GReinSS and GEM-based approaches are trained until Pr(X1 , . . . , XN | θ) converges. Naive policy gradients are trained until the achieved reward converges, and GFlowNets are trained until the loss function converges. The local search algorithm iterates through the observations X1 , . . . , XN and for each observation Xi it searches for the state S that maximizes Pr(Xi | S). For each Xi , local search starts with a binary vector representation of the latent state S and performs a sequence of modifications to this vector until reaching a locally optimal state S. Specifically, each modification selects the single element modification to S that maximizes Pr(Xi | S). This continues until convergence of Pr(Xi | S) such that no single element modification to S can improve Pr(Xi | S). For the set reconstruction simulations, local search is equivalent to including an element in the state S if and only if the value for that element in the observation is greater than 0.5. Consequently, local search is guaranteed to find the globally optimal S for maximizing Pr(Xi | S) for the set reconstruction simulations.
C. Additional Experimental Details C.1. Process Graph Inference Simulation Details In each graph-based simulation, latent states correspond to graphs that are generated by applying weight thresholding on some base graph. Each simulation instance has a randomly generated directed base graph with each directed edge being included with probability 1/2. Weights for each directed edge in the base graph are uniformly randomly generated from U ( 41 , 1) (with the lower limit set above 0 to ensure all edges in the base graph actually occur in some latent graph). 18
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
Each of the 1000 graphs in a simulation is generated by uniformly randomly generating a threshold from U (0, 1) and then performing weight thresholding (Yan et al., 2018) on the base graph with that threshold. Observations consist of the start and end points of random walks in the graph. Specifically, the walk begins by uniformly randomly selecting the start node. Then at each step we randomly choose between each edge directed out of that node as well as the “stop” action of ending on the current node. The stopping action and each outgoing edge are chosen with equal probability. Note that this implies if the current node has no outgoing edges then it will always be stopped on if reached. Each Xi consists of the list of start and end points of all of the random walks generated. The number of random walks generated is set to 10, 100, and 1000 in the three simulations. C.2. F1 Score Calculation The F1 score is defined as follows
F1 = 2
precision · recall 2 · true positives = . precision + recall (2 · true positives) + false positives + false negatives
(51)
Let Ŝi be a predicted state and Si∗ be the true latent state, where latent states represent either sets or graphs as in our simulations (Section 4.1). Then, true positives consist of elements or edges in both Ŝi and Si∗ . False positives consist of elements or edges in Ŝi but not Si∗ . False negatives consist of elements or edges in Si∗ but not Ŝi . From these three quantities, an F1 score is calculated for any predicted latent state Ŝi given the true latent state Si∗ . The distribution of F1 scores for the N predictions Ŝ1:N is either summarized by median values as in Figures 2 and 3, or visualized fully as in Figures S3, S5 and S6. C.3. Set Reconstruction Simulation Details The latent states correspond to sets of elements represented by binary vectors of size vsize . The default size of the vectors is 100, but additional modified simulations use vectors of size 10 and 1000. The sets are generated by taking the union of subsets of elements in a dictionary of reusable modules (Mairal et al., 2009). The number of subsets included in each state √ and the number of elements in each subset are both set to scale with vsize so that the proportion of elements included in √ each state remains roughly constant. Specifically, the dictionary contains vsize (rounded down) subsets/modules. Each √ subset is generated by randomly including each element with probability 2/ vsize . Each state is generated by randomly including each subset in the dictionary with probability 0.1. As an example with vsize = 10, if the dictionary contains the two subsets {2, 3, 5, 7} and {1, 2, 3, 4}, and the state includes both of these subsets, then the state would consist of {1, 2, 3, 4, 5, 7} which is equivalent to the binary vector [1, 1, 1, 1, 1, 0, 1, 0, 0, 0]. The observation would then consist of adding Gaussian noise with mean 0 and standard deviation σ to the vector [1, 1, 1, 1, 1, 0, 1, 0, 0, 0]. C.4. Off-policy Sampling Procedure for Simulated Set Reconstruction For generating each trajectory τ , the off-policy sampling procedure first, with 50% probability, either samples a trajectory on-policy or selects a random observation Xi from the set of observations X1 , . . . , XN . If an observation Xi is selected, we use the following procedure. Let Xi,1 , . . . , Xi,M ∈ R represent the observed values for each of the M elements in the universe. For any state S, let Sj ∈ {0, 1} represent whether or not the jth element is included in the state. Intuitively, if Xi,j is large (> 0.5) we wish to increase the probability that Sj = 1, and if Xi,j is small (< 0.5) we wish to decrease the probability that Sj = 1. The value σ is the standard deviation of the Gaussian noise in the simulation. Define PG (i, j) as the 2 log probability ratio of Sj taking the value 1 rather than the value 0 given Xi,j , which is equal to 2σ1 2 Xi,j − (1 − Xi,j )2 since Xi,j comes from a Gaussian with standard deviation σ centered at either 0 or 1. Given the state S and model parameters θ let PO (S, j, θ) be the logits proportional to the log probability of the next action being adding the element j to the set S. The new off-policy logits (proportional to the log probability of the next action adding state j to set S) are then PO (S, j, θ) + PG (i, j). The softmax operation then converts these logits to probabilities. Intuitively, the probability of adding the element j to the set S increases proportionally to PG (i, j), which is the log probability ratio that Sj = 1 rather than Sj = 0 given Xi,j . 19
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
C.5. Generative Model for RNA Splicing Junctions are defined as the connection between distinct exons in an isoform and can be represented by tuples containing the ending genomic position of one exon and the starting genomic position of the following exon. For instance, if an isoform contains the exon ranging from the genomic position 1000 to 1100 followed by the exon ranging from genomic position 2000 to 2100, then the junction connecting these exons can be represented by (1100, 2000). We choose to represent isoforms in terms of sequences of junction tuples. Note that the GTEx read count data with stable genomic positions is only publicly available in terms of the number of reads covering each junction. The exact nucleotides in an isoform can be precisely defined in terms of a sequence of junctions, with the minor exception of the exact length of the first and last exons. We model isoforms via a generative process, starting with an empty isoform and iteratively selecting junctions to include in the isoform. Since human data on GTEx has known possible exons, we allow our model to select its next junction j if the start position of this junction j and the end position of the previously selected junction j ′ define a known exon. For species without known exons, a reasonable approximation of this is to allow the model to select its next junction j if the start position of j is within a certain number of nucleotides (corresponding to a reasonable exon length) of the end of the previous junction. The first junction is allowed to be any junction starting in the end position of any starting exon in the GENCODE v39 annotation. The isoform is allowed to end on a junction if the junction’s ending position is the start position of any ending exon in the GENCODE v39 annotation. For species without known exons, one could approximate by allowing starting and ending positions within a certain distance of the lowest and highest genomic positions observed in the reads mapped to the relevant gene. The probability of each of the possible junctions at each step is determined by a neural network from an input of a one-hot encoding of the junctions already selected in the isoform. The neural network is specifically a two-layer neural network with 50 hidden neurons, and a leakyReLU non-linearity. To avoid ambiguity, we use the term sequencing sample to refer to a specific sequencing aliquot on GTEx, which is the unit with individual read count data. This is distinct from tissue samples from which one can derive multiple sequencing aliquots, and distinct from our policy learning sampling procedure. After generating the isoform, we select the probability of this isoform for each sequencing sample in the dataset, and the probability of observing reads from each junction in the isoform. Selecting the sequencing sample is performed using neural network with the same architecture as is used for constructing the isoform. Selecting the probability of reads from each junction is done by applying a learned sequencing bias vector with one log probability for each junction. The logit for each junction in the isoform is set to this sequencing bias value, and the logit for each junction not in the isoform is set to approximately negative infinity (technically, the bias minus 500, corresponding to a probability below e−500 to avoid numerical issues). The log probability of each junction is then calculated via a log softmax. Although isoforms are sampled according to the policy, off-policy sampling is used for selecting the sequencing sample and junction that the read comes from. Specifically, we simultaneously generate all possible sequencing samples and junctions in order to reduce reward variance during training. We therefore use the importance sampling correction term in our loss function. C.6. RNA Splicing Error Calculation We calculate the isoform prediction error using the Jaccard distance between isoforms as well as optimal transport to account for isoform proportions. This section will first describe the general calculation together with an illustrative example, before moving on to technical details for the GTEx dataset. The Jaccard index for measuring similarity between two sets is defined as one minus the size of the intersection of the two sets divided by the size of the union of the two sets. The Jaccard distance is then one minus the Jaccard index. For genomic regions, the Jaccard distance is then one minus the ratio of the number of nucleotides in the intersection of the two regions to the number of nucleotides in the union of the two regions. A hypothetical example is shown in Figure S7a. In this example, isoform 1 consists of an exon from 100bp to 200bp, an exon from 300bp to 400bp, and an exon from 500bp to 550bp, while isoform 2 consists of an exon from 150bp to 250bp and an exon from 300bp to 400bp. Then, the intersection is 150bp to 200bp as well as 300bp to 400bp, totaling 150 nucleotides. The union is 100bp to 250bp, 300bp to 400bp, and 500bp to 550bp, totaling 300bp. Thus, the Jaccard distance in this example is 1 − (150/300) = 0.5. For each tissue sample, there can be multiple isoforms with different proportions detected by long-read sequencing, and multiple isoforms with different proportions predicted from short-read sequencing by either GReinSS or RSEM. Figure S7b shows a hypothetical example where long-read sequencing detects 80% proportion to isoform 1 and 20% proportion to isoform 2, whereas the prediction assigns 60% proportion to isoform 1 and 40% proportion to isoform 2. To account for the proportions of each isoform on each sample, we calculate a Jaccard distance matrix between all pairs of predicted isoforms and long-read detected isoforms, and then apply optimal transport to these proportions. This is 20
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
a
b
d
c
Figure S2. Ablation of policy learning methods on the set reconstruction simulation to remove off-policy sampling. a-b Policy learning methods consistently achieve a higher performance when off-policy sampling is used. Panel a has varying standard deviations σ and panel b has varying sizes of the universe of elements. c-d Off-policy GReinSS consistently outperforms all off-policy and non-RL machine learning baselines. Panel c has varying standard deviations σ and panel d has varying sizes of the universe of elements.
illustrated for our hypothetical example in Figure S7c. Specifically, the distance matrix contains 0.5 as the distance between isoform 1 and isoform 2, and contains the distance 0 between each isoform and itself. Optimal transport then assigns 20% proportion to the prediction and long-read agreeing on isoform 1, 60% proportion to the prediction and long-read agreeing on isoform 2, and 20% proportion to the prediction having isoform 2 but long-read having isoform 1. The 20% proportion of the prediction having isoform 2 but long-read having isoform 1 results in an error of 20% of the Jaccard distance between isoform 1 and isoform 2, namely 0.1 = 0.5 · 20%. In general, optimal transport yields a correspondence between predicted isoforms and long-read isoforms that minimizes the error. As a special case, if there is only one long-read detected isoform and only one predicted isoform, optimal transport simplifies to the Jaccard distance between these two isoforms. After calculating errors on each tissue sample with optimal transport, we calculate a weighted average of errors across tissue samples. Specifically, each sample is weighted by the long-read sequencing read count (of reads covering whole isoforms). This weighting is essential since some tissue samples have very few or even zero long-read sequencing reads (covering whole isoforms). Additional details: Long read sequencing directly detects entire isoforms through reads that cover entire isoforms. Consequently, we filter for reads covering entire isoforms when calculating long-read sequencing read counts. As another technical detail, the GTEx database can have multiple sequencing aliquots for each tissue sample that have their own RNA sequencing data. To compare isoform estimations between short-read sequencing data and long-read sequencing data, we first pool the predicted isoform quantities across sequencing aliquots for each tissue sample. Then, for the same set of 61 tissue samples, we have isoforms predicted from short-read sequencing data from GReinSS and RSEM, as well as isoforms quantified from long-read sequencing using FLAIR. As discussed in Section C.5, we define isoforms in terms of their list of junctions. Therefore, we define the Jaccard distance with the region of an isoform defined as including all exons between the first and last junctions (i.e., only excluding the starting and ending exons themselves). We then calculate the Jaccard distance between all pairs of predicted isoforms and the isoforms estimated from long-read sequencing to form a distance matrix used in the optimal transport error calculation.
21
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
10 random walks
100 random walks
1000 random walks
Figure S3. The distribution of F1 scores of each method on graph inference simulations with 10, 100 and 1000 random walks per observation. The difficulty increases for smaller numbers of random walks, resulting in a larger advantage for GReinSS.
1.0 0.9
F1 score
0.8
GReinSS GReinSS (50% more off-policy bias) GReinSS (50% less off-policy bias) VAE autoregressive diffusion local search naive policy learning GFlowNet
0.7 0.6 0.5 0.4 0.1
0.2
0.3 noise level
0.4
0.5
Figure S4. GReinSS is not sensitive to changes in the off-policy sampling proposal. We tested GReinSS’s sensitivity to modifying the off-policy sampling proposal by artificially increasing and decreasing the strength of the sampling bias on the action logits by 50%. There was very little impact on the F1 score (with a maximum decrease of 0.0044 for any σ value), with GReinSS remaining the top performing method.
22
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
𝜎 = 0.1
𝜎 = 0.4
𝜎 = 0.3
𝜎 = 0.2
𝜎 = 0.5
Figure S5. The distribution of F1 scores of each method on set inference simulations with standard deviations σ set to 0.1, 0.2, 0.3, 0.4, and 0.5. The difficulty increases as the variance increases, but GReinSS maintains the top performance. universe size 1000
universe size 100
universe size 10
Figure S6. The distribution of F1 scores of each method on set inference simulations with the size of the universe of elements set to 10, 100 and 1000. The difficulty increases as the universe of elements increases in size, but GReinSS maintains the top performance.
a
b
junction isoform 1:
exon
isoform 2:
exon
exon junction
junction
long read: exon
isoform 1:
exon
exon isoform 2:
reference gene
isoform 1:
exon
exon
c
exon
80% distance isoform 1: isoform 2: matrix: exon
20%
300bp union: 150bp overlap:
Jaccard distance = 1 – (150bp/300bp) = 0.5
predicted: exon
isoform 2:
exon exon
exon
60% exon
40%
isoform 1: isoform 2: 0.5 0 0 0.5
optimal transport: true proportion: predicted 20% 80% isoform 1: isoform 2: proportion: 60% isoform 1: 60% 0% 20% 20% 40% isoform 2: error = 0.5 ⋅ 20% = 0.1
Figure S7. A hypothetical example of calculating the isoform prediction error. a Isoform 1 consists of an exon from 100bp to 200bp and an exon from 350bp to 450bp. Isoform 2 consists of an exon from 100bp to 200bp, an exon from 300bp to 400bp, and an exon from 500bp to 550bp. The intersection is 150bp, and the union is 300bp, resulting in a Jaccard distance of 1 − (150/300) = 0.5. b In this example, the long-read sequencing has isoform 1 with proportion 80% and isoform 2 with proportion 20%, while the prediction has isoform 1 with proportion 60% and isoform 2 with proportion 40%. c The error is calculated using optimal transport. Specifically, the prediction has 20% more of isoform 2 that should be predicted as isoform 1, resulting in an error of 20% of the Jaccard distance of 0.5 between isoform 1 and 2, namely 0.1.
23
Generative Modeling of Discrete Latent Structures via Dynamic Policy Gradients
isoform prediction error
1.0 0.8 0.6 0.4 0.2 0.0
GReinSS
RSEM
GFlowNet
naive policy gradients
Figure S8. GReinSS outperforms GFlowNets and naive policy gradients in the isoform reconstruction task. GFlowNets and naive policy gradients were run on 100 random genes for the RNA isoform reconstruction task. For these genes, the median isoform prediction error is 0.141, 0.317, 0.388, and 0.427 for GReinSS, RSEM, GFlowNets, and naive policy gradients, respectively. The poor performance of GFlowNets and naive policy gradients clearly indicates that a neural network reparameterization of isoform generation is insufficient to achieve high accuracy, and the GReinSS approach is necessary to achieve this improvement.
a
b
Figure S9. Additional comparisons of GReinSS and RSEM errors in isoform proportion prediction. a The full boxplot of the GReinSS error minus the RSEM error, including all outliers. b A scatterplot showing the GReinSS error and RSEM error for all genes.
24