Probabilistic Inference in Language Models via Twisted Sequential Monte Carlo

Stephen ZhaoRob BrekelmansAlireza MakhzaniRoger Baker Grosse

article2024ICML90 citationsBest Paper Award

Proposes a twisted Sequential Monte Carlo framework using contrastively learned lookahead functions to steer language model generation toward sequence-level targets and establish bidirectional bounds for evaluating inference quality.

Listen

Steering large language models to satisfy specific capability, alignment, and safety requirements is increasingly vital for real-world deployment. Common objectives—such as reinforcement learning from human feedback, automated red-teaming, prompt optimization, and infilling—can be unified mathematically as sampling from a non-causal, unnormalized target distribution. In these settings, standard generation strategies often prove inefficient or exhibit mode-dropping, as they evaluate full-sequence reward functions only after the complete output has already been generated.

The article establishes a probabilistic inference framework using twisted Sequential Monte Carlo to steer text generation and rigorously evaluate the accuracy of language model sampling. It evaluates how learned predictive functions—termed twist functions—can estimate the future expected value of a terminal reward to guide intermediate token generation. The approach is examined across multiple tasks, including rare toxic text elicitation, sentiment-controlled review writing, and text infilling, utilizing both small language models and medium-sized architectures.

To guide inference efficiently, the article introduces Contrastive Twist Learning, an energy-based approach that trains twist functions by matching intermediate model marginals to true target marginals across time. For model evaluation, the article proposes a bidirectional bounding method on the normalization constant (log partition function) using extended-state-space Sequential Monte Carlo. These bounds allow stakeholders to compute tight upper and lower bounds on target normalization constants and quantify the exact directional divergence between inference proposals and desired distributions, even for rare outcomes.

The core findings reveal distinct trade-offs across inference techniques: First, combining Sequential Monte Carlo resampling with contrastive twist-induced proposals yields near-exact target sampling with orders of magnitude fewer samples compared to standard importance sampling. Second, Contrastive Twist Learning consistently achieves balanced, superior performance across directional divergence metrics on toxic story and sentiment tasks, avoiding the severe mode-dropping observed in standard reinforcement learning methods like Proximal Policy Optimization. Third, while standard policy gradients minimize one-sided divergence effectively, maximum-likelihood-style distributional policy gradients excel when exact target data are readily accessible, reducing distributional mismatch significantly in infilling tasks.

These results carry significant operational implications for AI safety, compliance, and inference cost. By using learned twists to prune unpromising trajectories mid-generation, teams can efficiently discover rare failure modes (such as hidden toxicity or adversarial vulnerabilities) without wasteful brute-force generation. Furthermore, the bidirectional bounds provide a reliable verification mechanism to assess mode collapse and distributional coverage in fine-tuned models, reducing the risk of deploying over-optimized but brittle systems.

Organizations developing or deploying aligned models should adopt Sequential Monte Carlo bounds to audit generation quality and mode retention. When full target samples are unavailable, Contrastive Twist Learning is recommended to train guided decoders; when exact samples exist, distributional policy gradient approaches should be prioritized. Key limitations include the requirement of exact reference samples to calculate upper partition bounds and the computational overhead of training auxiliary twist networks. However, for applications where safety verification and controlled generation are critical, the methodology provides high confidence and substantial efficiency gains.

Cover for Probabilistic Inference in Language Models via Twisted Sequential Monte Carlo

Abstract

Numerous capability and safety techniques of Large Language Models (LLMs), including RLHF, automated red-teaming, prompt engineering, and infilling, can be cast as sampling from an unnormalized target distribution defined by a given reward or potential function over the full sequence. In this work, we leverage the rich toolkit of Sequential Monte Carlo (SMC) for these probabilistic inference problems. In particular, we use learned twist functions to estimate the expected future value of the potential at each timestep, which enables us to focus inference-time computation on promising partial sequences. We propose a novel contrastive method for learning the twist functions, and establish connections with the rich literature of soft reinforcement learning. As a complementary application of our twisted SMC framework, we present methods for evaluating the accuracy of language model inference techniques using novel bidirectional SMC bounds on the log partition function. These bounds can be used to estimate the KL divergence between the inference and target distributions in both directions. We apply our inference evaluation techniques to show that twisted SMC is effective for sampling undesirable outputs from a pretrained model (a useful component of harmlessness training and automated red-teaming), generating reviews with varied sentiment, and performing infilling tasks.

Table of Contents

  • 1. Introduction
  • 2. Background
  • 2.1. Simple Importance Sampling
  • 2.2. Sequential Monte Carlo
  • 3. Twisted Sequential Monte Carlo for Language Modeling
  • 3.1. Twist Functions
  • 3.2. Proposal Distribution
  • 3.3. Conditional Target Distributions
  • 3.4. Connections with Reinforcement Learning
  • 4. Learning the Twist Functions
  • 4.1. Contrastive Twist Learning
  • 4.1.1. Approximate Negative Sampling
  • 4.1.2. (Approximate) Positive Sampling
  • 4.2. Twist Learning Methods from Related Work
  • 5. Evaluating Inference in Language Models
  • 5.1. Applications of log Z σ Estimation
  • 5.2. Bidirectional SMC Bounds on log Z σ
  • 6. Related Work
  • 7. Experiments
  • 7.1. Comparing SIS and SMC for log Z σ Estimation
  • 7.2. Evaluating Twist-Induced or Variational Proposals
  • 7.2.1. Generating Toxic Stories
  • 7.2.2. Generation with Varied Sentiment
  • 7.2.3. Infilling
  • 8. Conclusion
  • Acknowledgments
  • Impact Statement
  • References
  • A. Proofs
  • A.1. Proof for Optimal Intermediate Target Distributions
  • A.2. Proof of Twist-Induced Proposal
  • A.3. Derivation of CTL Gradient
  • B. SMC with Intermediate Potentials and Connection with Soft Reinforcement Learning
  • B.1. Twisted SMC with Intermediate Potentials
  • B.2. Conditional Twisted SMC
  • B.3. Connection with Soft Reinforcement Learning
  • B.4. Remarks on Parameterization
  • C. Twist Learning Losses
  • C.1. Soft Q-Learning (RL) and Path Consistency Losses from Log Importance Weights
  • C.1.1. Soft Q-Learning and RL Baseline
  • C.1.2. Path Consistency Learning (for Twist Learning)
  • C.2. Controlled Decoding Losses via Optimal Twist Identities (Mudgal et al., 2023)
  • C.3. SIXO: Smoothing Inference with Twisted Objectives (Lawson et al., 2022)
  • C.4. FUDGE: Future Discriminators (Yang & Klein, 2021)
  • D. Decoding Strategies using Learned Twists from Mudgal et al. (2023)
  • D.1. Proposal Sampling in Mudgal et al. (2023)
  • D.2. Blockwise Greedy Decoding in Mudgal et al. (2023)
  • E. Proposal Learning Methods
  • E.1. Path Consistency Learning for Controlled Generation
  • E.2. Policy Gradient Methods
  • E.3. Policy Gradient with Mass-Covering / Maximum Likelihood KL Divergence
  • E.3.1. Naive Use of Proposal Learning to Define Twisted SMC Targets
  • F. Bidirectional SMC
  • G. Additional Experiment Details
  • G.1. Common Details Across Experiments
  • G.2. Choices of Twist Parameterization
  • G.2.1. Linear Head
  • G.2.2. MLP Head
  • G.2.3. Separate Transformer for the Twist
  • G.2.4. Separate Transformer for the Twist, with MLP Head
  • G.3. Comments on Our Choices of Experiment Settings
  • G.4. Experiment-Specific Details
  • H. Additional Experimental Results
  • H.1. Qualitative Results
  • H.2. Infilling with Fewer Tokens
  • H.3. Approximate vs. Exact Posterior Sampling

Knowls

  1. Knowl 1 — Language-model steering as inference under a terminal potential

    definition

    For a fixed prompt s0s_0, let s1:Ts_{1:T} be a generated token sequence, p0(s1:T∣s0)p_0(s_{1:T}\mid s_0) the pretrained language model, and ϕ(s1:T)≥0\phi(s_{1:T})\geq 0 a potential scoring complete sequences. The target distribution is

    σ(s1:T∣s0)=p0(s1:T∣s0)ϕ(s1:T)Zσ(s0),Zσ(s0)=∑s1:Tp0(s1:T∣s0)ϕ(s1:T).\sigma(s_{1:T}\mid s_0)=\frac{p_0(s_{1:T}\mid s_0)\phi(s_{1:T})}{Z_\sigma(s_0)},\qquad Z_\sigma(s_0)=\sum_{s_{1:T}}p_0(s_{1:T}\mid s_0)\phi(s_{1:T}).

    The potential can encode a reward, classifier probability, verifier score, or indicator of a desired sequence property. The partition function is generally intractable because it sums over all complete sequences. The paper treats controlled generation, including steering toward undesirable outputs, as sampling from this target.

  2. Knowl 2 — Twisted sequential Monte Carlo for terminal-potential targets

    model/method

    Twisted SMC approximates sampling from a complete-sequence target by maintaining KK partial sequences and resampling promising prefixes as generation proceeds. At step tt, particle kk samples a token from a proposal q(st∣s1:t−1k)q(s_t\mid s^k_{1:t-1}) and appends it to its prefix. Let ψt(s1:t)\psi_t(s_{1:t}) be a nonnegative twist, set ψ0=1\psi_0=1, and use ϕ\phi as the terminal potential. The incremental weights are

    wtk={p0(stk∣s1:t−1k)q(stk∣s1:t−1k)ψt(s1:tk)ψt−1(s1:t−1k),t<T,p0(sTk∣s1:T−1k)q(sTk∣s1:T−1k)ϕ(s1:Tk)ψT−1(s1:T−1k),t=T.w_t^k=\begin{cases} \dfrac{p_0(s_t^k\mid s^k_{1:t-1})}{q(s_t^k\mid s^k_{1:t-1})}\dfrac{\psi_t(s^k_{1:t})}{\psi_{t-1}(s^k_{1:t-1})},&t<T,\\[6pt] \dfrac{p_0(s_T^k\mid s^k_{1:T-1})}{q(s_T^k\mid s^k_{1:T-1})}\dfrac{\phi(s^k_{1:T})}{\psi_{T-1}(s^k_{1:T-1})},&t=T. \end{cases}

    At each intermediate step, resample the KK prefixes with replacement according to their normalized incremental weights, then continue generation from the selected prefixes. With resampling at each t<Tt<T, the partition-function estimate is Z^σ=∏t=1T(K−1∑k=1Kwtk)\widehat Z_\sigma=\prod_{t=1}^T\bigl(K^{-1}\sum_{k=1}^K w_t^k\bigr). Using the exact terminal potential in the final weight makes this estimator unbiased for ZσZ_\sigma, regardless of the intermediate twists and proposal. Resampling can also be scheduled at selected or adaptive steps.

  3. Knowl 3 — Optimal twists are future expected target potentials

    theoretical result

    For a target σ(s1:T)∝p0(s1:T)ϕ(s1:T)\sigma(s_{1:T})\propto p_0(s_{1:T})\phi(s_{1:T}), the optimal twist at a prefix s1:ts_{1:t} is, up to a positive constant independent of that prefix, the expected terminal potential under future continuations of the base model:

    ψt∗(s1:t)∝∑st+1:Tp0(st+1:T∣s1:t)ϕ(s1:T),t<T.\psi_t^*(s_{1:t})\propto\sum_{s_{t+1:T}}p_0(s_{t+1:T}\mid s_{1:t})\phi(s_{1:T}),\qquad t<T.

    With intermediate targets formed as p0(s1:t)ψt(s1:t)p_0(s_{1:t})\psi_t(s_{1:t}) and normalized, these optimal twists make every intermediate target equal the true target marginal σ(s1:t)\sigma(s_{1:t}). They also obey the recursion ψt∗(s1:t)∝∑st+1p0(st+1∣s1:t)ψt+1∗(s1:t+1)\psi_t^*(s_{1:t})\propto\sum_{s_{t+1}}p_0(s_{t+1}\mid s_{1:t})\psi_{t+1}^*(s_{1:t+1}), with the terminal twist given by ϕ\phi. These twists are generally unavailable because computing them requires summing over future continuations.

  4. Knowl 4 — Twists induce a variance-minimizing one-step proposal

    theoretical result

    Given an intermediate target πt(s1:t)\pi_t(s_{1:t}) and its preceding target πt−1(s1:t−1)\pi_{t-1}(s_{1:t-1}), the proposal minimizing the variance of the one-step incremental importance weight is the conditional target transition:

    qtπ(st∣s1:t−1)=πt(s1:t)∑st′πt(s1:t−1,st′)=p0(st∣s1:t−1)ψt(s1:t)∑st′p0(st′∣s1:t−1)ψt(s1:t−1,st′).q_t^\pi(s_t\mid s_{1:t-1})=\frac{\pi_t(s_{1:t})}{\sum_{s_t'}\pi_t(s_{1:t-1},s_t')}=\frac{p_0(s_t\mid s_{1:t-1})\psi_t(s_{1:t})}{\sum_{s_t'}p_0(s_t'\mid s_{1:t-1})\psi_t(s_{1:t-1},s_t')}.

    Here, sts_t ranges over the model vocabulary and the twist defines the intermediate target. For t<Tt<T, the proposal can be normalized by evaluating the twisted scores over next-token choices. At the final step, when evaluating the terminal potential over every possible next token is costly, the paper uses an approximate terminal twist to form the proposal and still evaluates the true ϕ\phi on sampled complete sequences for the final importance weights.

  5. Knowl 5 — Contrastive twist learning matches every target prefix marginal

    model/method

    Contrastive twist learning (CTL) trains parameterized twists ψtθ\psi_t^\theta by matching their normalized intermediate targets πtθ\pi_t^\theta to the target sequence marginals. Its objective is

    LCTL(θ)=∑t=1TDKL ⁣(σ(s1:t) ∥ πtθ(s1:t)),πtθ(s1:t)∝p0(s1:t)ψtθ(s1:t)(t<T),\mathcal L_{\mathrm{CTL}}(\theta)=\sum_{t=1}^T D_{\mathrm{KL}}\!\left(\sigma(s_{1:t})\,\middle\|\,\pi_t^\theta(s_{1:t})\right),\qquad \pi_t^\theta(s_{1:t})\propto p_0(s_{1:t})\psi_t^\theta(s_{1:t})\quad(t<T),

    with the terminal target fixed by p0(s1:T)ϕ(s1:T)p_0(s_{1:T})\phi(s_{1:T}). For each timestep, the negative gradient is the target-marginal expectation of ∇θlog⁡ψtθ\nabla_\theta\log\psi_t^\theta minus its expectation under πtθ\pi_t^\theta. Thus learning contrasts positive target prefixes with negative prefixes from the current twisted target. Positive prefixes can come from exact target samples, rejection sampling, or importance-weighted SIS/SMC samples; negative expectations can be estimated with SIS or SMC. Exact complete-sequence samples are truncated to obtain prefixes, while approximate positive sampling uses the final target weights. Matching with the forward KL is intended to be mass-covering, reducing the risk that an early twist prunes prefixes that later lead to high target probability. A practical limitation for thresholded targets is that if no sampled sequence passes the threshold, all positive importance weights can vanish and provide no learning signal.

  6. Knowl 6 — Bidirectional SMC bounds sandwich the log partition function

    theoretical result

    Let SS denote the extended random state containing the particle tokens and resampling indices of a KK-particle SMC run. Let qSMC(S)q_{\mathrm{SMC}}(S) be the law of the ordinary SMC procedure and σSMC(S)\sigma_{\mathrm{SMC}}(S) the corresponding extended target law, constructed by including an exact sample from the sequence target. If wtiw_t^i is the incremental importance weight of particle ii at step tt, define

    W(S)=∏t=1T(1K∑i=1Kwti).W(S)=\prod_{t=1}^T\left(\frac{1}{K}\sum_{i=1}^K w_t^i\right).

    Then the log partition function obeys

    EqSMC[log⁡W(S)]≤log⁡Zσ≤EσSMC[log⁡W(S)].\mathbb E_{q_{\mathrm{SMC}}}[\log W(S)]\leq\log Z_\sigma\leq\mathbb E_{\sigma_{\mathrm{SMC}}}[\log W(S)].

    The lower-bound gap is DKL(qSMC∥σSMC)D_{\mathrm{KL}}(q_{\mathrm{SMC}}\|\sigma_{\mathrm{SMC}}) and the upper-bound gap is DKL(σSMC∥qSMC)D_{\mathrm{KL}}(\sigma_{\mathrm{SMC}}\|q_{\mathrm{SMC}}). The upper-bound expectation requires an exact target sequence: the sampling procedure keeps its lineage represented at each resampling step, resamples the other particles using the importance weights, and samples proposals for the remaining particles. Both bounds become exact as KK grows; with no intermediate resampling, they reduce to the importance-weighted sampling bounds.

  7. Knowl 7 — Partition-function bounds evaluate inference in both KL directions

    model/method

    For a tractable proposal distribution q(s1:T)q(s_{1:T}) and target σ(s1:T)=p0(s1:T)ϕ(s1:T)/Zσ\sigma(s_{1:T})=p_0(s_{1:T})\phi(s_{1:T})/Z_\sigma, the paper estimates divergence in both directions using a log-partition estimate and samples from the relevant distributions:

    DKL(q∥σ)=Eq ⁣[log⁡q(s1:T)p0(s1:T)ϕ(s1:T)]+log⁡Zσ,D_{\mathrm{KL}}(q\|\sigma)=\mathbb E_q\!\left[\log\frac{q(s_{1:T})}{p_0(s_{1:T})\phi(s_{1:T})}\right]+\log Z_\sigma, DKL(σ∥q)=Eσ ⁣[log⁡p0(s1:T)ϕ(s1:T)q(s1:T)]−log⁡Zσ.D_{\mathrm{KL}}(\sigma\|q)=\mathbb E_\sigma\!\left[\log\frac{p_0(s_{1:T})\phi(s_{1:T})}{q(s_{1:T})}\right]-\log Z_\sigma.

    The reverse divergence additionally requires target samples. For a single output returned by a KK-sample SIS or SMC procedure, the particle-set marginal is generally intractable, but the bidirectional bounds still diagnose sample quality: by data processing, the extended-space KL gaps bound the corresponding KL divergences between the returned-sample distribution and σ\sigma. A small gap between the log-partition upper and lower bounds therefore indicates that the sampling procedure is close to the target in symmetrized KL. Evaluating both directions is useful because a proposal can have low DKL(q∥σ)D_{\mathrm{KL}}(q\|\sigma) while still dropping target modes, which is exposed by DKL(σ∥q)D_{\mathrm{KL}}(\sigma\|q).

  8. Knowl 8 — Conditional twists yield exact posterior samples for infilling

    model/method

    For an observation oTo_T with likelihood σ(oT∣s1:T)\sigma(o_T\mid s_{1:T}), conditional twisted SMC targets the posterior

    σ(s1:T∣oT)=p0(s1:T)σ(oT∣s1:T)Zσ(oT),Zσ(oT)=∑s1:Tp0(s1:T)σ(oT∣s1:T).\sigma(s_{1:T}\mid o_T)=\frac{p_0(s_{1:T})\sigma(o_T\mid s_{1:T})}{Z_\sigma(o_T)},\qquad Z_\sigma(o_T)=\sum_{s_{1:T}}p_0(s_{1:T})\sigma(o_T\mid s_{1:T}).

    The optimal twist at prefix s1:ts_{1:t} is proportional to the future observation likelihood σ(oT∣s1:t)\sigma(o_T\mid s_{1:t}), which marginalizes over continuations of the prefix. The paper uses a conditional twist network taking both the prefix and observation as inputs. For infilling, oTo_T is a continuation of cc tokens, with likelihood given by the base model's probability of that continuation after the candidate prefix. An exact posterior sample for a sampled continuation is obtained by drawing the full prefix-plus-continuation sequence from the base model and treating its prefix as a sample from the posterior conditioned on its generated continuation. This construction supplies exact target samples for training and for the upper partition-function bound.

  9. Knowl 9 — Twisted SMC recovers soft reinforcement learning values

    model/method

    For a terminal reward r(s1:T)r(s_{1:T}) and regularization strength β>0\beta>0, setting ϕ(s1:T)=exp⁡(βr(s1:T))\phi(s_{1:T})=\exp(\beta r(s_{1:T})) makes the target proportional to p0(s1:T)exp⁡(βr(s1:T))p_0(s_{1:T})\exp(\beta r(s_{1:T})). In this case, an optimal twist is an exponentiated soft action value: ψt∗(s1:t)=exp⁡(βQt∗(st,s1:t−1))\psi_t^*(s_{1:t})=\exp(\beta Q_t^*(s_t,s_{1:t-1})), up to prefix-independent scaling. With no intermediate reward, the optimal values obey the soft Bellman recursion

    Qt∗(st,s1:t−1)=1βlog⁡∑st+1p0(st+1∣s1:t)exp⁡ ⁣(βQt+1∗(st+1,s1:t)).Q_t^*(s_t,s_{1:t-1})=\frac{1}{\beta}\log\sum_{s_{t+1}}p_0(s_{t+1}\mid s_{1:t})\exp\!\left(\beta Q_{t+1}^*(s_{t+1},s_{1:t})\right).

    The corresponding twist-induced proposal is proportional to p0(st∣s1:t−1)exp⁡(βQt∗(st,s1:t−1))p_0(s_t\mid s_{1:t-1})\exp(\beta Q_t^*(s_t,s_{1:t-1})). This gives the twists the role of a soft-RL critic and the proposal the role of an actor. The same target is the optimizer of expected reward minus β−1DKL(q∥p0)\beta^{-1}D_{\mathrm{KL}}(q\|p_0) over sequence distributions qq.

  10. Knowl 10 — Experiments show complementary strengths across toxicity, sentiment, and infilling

    data/table

    The evaluation compares the two directional KL divergences for proposals from twisted CTL, soft-RL twist learning, SIXO, and FUDGE, and for the direct proposals DPG and PPO. Toxicity uses TinyStories with a classifier-probability target and 20 generated tokens; sentiment uses GPT-2 Medium, the prompt “I bought this,” a one-star classifier target, and 10 generated tokens. Infilling uses TinyStories with 15 generated tokens conditioned on 10 continuation tokens, and reports KL values averaged over 2,000 continuations. Reported values are means with 95% confidence intervals over five training seeds; infilling divergences are averaged over conditioning continuations.

    Task Method DKL(q∥σ)D_{\mathrm{KL}}(q\|\sigma) DKL(σ∥q)D_{\mathrm{KL}}(\sigma\|q)
    Toxicity Twisted Contrastive 1.11±0.051.11\pm0.05 1.07±0.021.07\pm0.02
    Toxicity Twisted RL 1.52±0.091.52\pm0.09 1.42±0.031.42\pm0.03
    Toxicity Twisted SIXO 1.71±0.061.71\pm0.06 1.98±0.041.98\pm0.04
    Toxicity Twisted FUDGE 3.24±0.263.24\pm0.26 2.00±0.132.00\pm0.13
    Toxicity DPG 1.09±0.051.09\pm0.05 1.12±0.031.12\pm0.03
    Toxicity PPO 0.98±0.010.98\pm0.01 1.32±0.041.32\pm0.04
    Sentiment Twisted Contrastive 0.55±0.030.55\pm0.03 0.47±0.010.47\pm0.01
    Sentiment Twisted RL 0.94±0.040.94\pm0.04 0.81±0.020.81\pm0.02
    Sentiment Twisted SIXO 0.73±0.030.73\pm0.03 0.59±0.020.59\pm0.02
    Sentiment Twisted FUDGE 1.01±0.071.01\pm0.07 0.77±0.070.77\pm0.07
    Sentiment DPG 0.72±0.040.72\pm0.04 0.57±0.010.57\pm0.01
    Sentiment PPO 1.04±0.311.04\pm0.31 0.87±0.200.87\pm0.20
    Infilling Twisted Contrastive 23.93±0.3423.93\pm0.34 8.87±0.058.87\pm0.05
    Infilling Twisted RL 31.35±2.3331.35\pm2.33 14.96±1.6914.96\pm1.69
    Infilling Twisted SIXO 20.34±0.3620.34\pm0.36 7.43±0.047.43\pm0.04
    Infilling Twisted FUDGE 60.93±2.8260.93\pm2.82 19.85±0.5119.85\pm0.51
    Infilling DPG 13.27±0.4413.27\pm0.44 4.90±0.034.90\pm0.03
    Infilling PPO 19.37±0.4119.37\pm0.41 14.07±0.5014.07\pm0.50

    CTL has the lowest reverse KL on toxicity and the lowest values in both directions on sentiment; PPO has the lowest forward KL on toxicity. DPG has the lowest values in both directions on infilling, where its use of exact posterior samples is advantageous. Thus, no method dominates every task or divergence. In a separate rare-event toxicity experiment, the target used an indicator requiring a non-toxic classifier logit at most −5-5 (corresponding to greater than 99% probability of toxicity). With the base-model proposal, SIS generally failed to find qualifying sequences, while SMC resampling with learned twists eventually produced tight bounds as particle count increased. A twist-induced proposal made both SIS and SMC bounds tight with orders-of-magnitude fewer particles than base-model sampling; ESS-based resampling gave similar results to resampling at every step.

Coverage note — Detailed derivations of alternative twist-learning baselines and supplementary hyperparameter and qualitative analyses are omitted because they support, rather than add to, the central inference framework, CTL method, bidirectional bounds, and main experimental comparisons.

References

  1. 1.Andrieu, C., Doucet, A., and Holenstein, R. Particle markov chain monte carlo methods. Journal of the Royal Statistical Society Series B: Statistical Methodology, 72(3): 269–342, 2010.
  2. 2.Anil, C., Zhang, G., Wu, Y., and Grosse, R. Learning to give checkable answers with prover-verifier games. arXiv preprint arXiv:2108.12099, 2021.
  3. 3.Bae, J., Zhang, M. R., Ruan, M., Wang, E., Hasegawa, S., Ba, J., and Grosse, R. B. Multi-rate vae: Train once, get the full rate-distortion curve. In The Eleventh International Conference on Learning Representations, 2022.
  4. 4.Bai, Y., Jones, A., Ndousse, K., Askell, A., Chen, A., Das-Sarma, N., Drain, D., Fort, S., Ganguli, D., Henighan, T., et al. Training a helpful and harmless assistant with reinforcement learning from human feedback. arXiv preprint arXiv:2204.05862, 2022.
  5. 5.Banerjee, A., Guo, X., and Wang, H. On the optimality of conditional expectation as a bregman predictor. IEEE Transactions on Information Theory, 51(7), 2005.
  6. 6.Brekelmans, R., Huang, S., Ghassemi, M., Ver Steeg, G., Grosse, R. B., and Makhzani, A. Improving mutual information estimation with annealed and energy-based bounds. In International Conference on Learning Representations, 2021.
  7. 7.Briers, M., Doucet, A., and Maskell, S. Smoothing algorithms for state–space models. Annals of the Institute of Statistical Mathematics, 62:61–89, 2010.
  8. 8.Burda, Y., Grosse, R., and Salakhutdinov, R. Importance weighted autoencoders. arXiv preprint arXiv:1509.00519, 2015.
  9. 9.Chopin, N., Papaspiliopoulos, O., et al. An introduction to sequential Monte Carlo, volume 4. Springer, 2020.
  10. 10.Cobbe, K., Kosaraju, V., Bavarian, M., Chen, M., Jun, H., Kaiser, L., Plappert, M., Tworek, J., Hilton, J., Nakano, R., et al. Training verifiers to solve math word problems. arXiv preprint arXiv:2110.14168, 2021.
  11. 11.Correa, N. K. Aira, 2023. URL https://huggingface.co/nicholasKluge/ToxicityModel.
  12. 12.Dathathri, S., Madotto, A., Lan, J., Hung, J., Frank, E., Molino, P., Yosinski, J., and Liu, R. Plug and play language models: A simple approach to controlled text generation. In International Conference on Learning Representations, 2019.
  13. 13.Del Moral, P., Doucet, A., and Jasra, A. Sequential monte carlo samplers. Journal of the Royal Statistical Society Series B: Statistical Methodology, 68(3):411–436, 2006.
  14. 14.Deng, H. and Raffel, C. Reward-augmented decoding: Efficient controlled text generation with a unidirectional reward model. In The 2023 Conference on Empirical Methods in Natural Language Processing, 2023.
  15. 15.Dohan, D., Xu, W., Lewkowycz, A., Austin, J., Bieber, D., Lopes, R. G., Wu, Y., Michalewski, H., Saurous, R. A., Sohl-Dickstein, J., et al. Language model cascades. arXiv preprint arXiv:2207.10342, 2022.
  16. 16.Domke, J. and Sheldon, D. R. Importance weighting and variational inference. Advances in neural information processing systems, 31, 2018.
  17. 17.Doucet, A., De Freitas, N., Gordon, N. J., et al. Sequential Monte Carlo methods in practice, volume 1. Springer, 2001.
  18. 18.Eikema, B., Kruszewski, G., Dance, C. R., Elsahar, H., and Dymetman, M. An approximate sampler for energy-based models with divergence diagnostics. Transactions on Machine Learning Research, 2022.
  19. 19.Eldan, R. and Li, Y. Tinystories: How small can language models be and still speak coherent english? arXiv preprint arXiv:2305.07759, 2023. URL https://huggingface.co/roneneldan/TinyStories-33M.
  20. 20.Finke, A. On extended state-space constructions for Monte Carlo methods. PhD thesis, University of Warwick, 2015.
  21. 21.Go, D., Korbak, T., Kruszewski, G., Rozen, J., Ryu, N., and Dymetman, M. Aligning foundation models for language with preferences through f-divergence minimization. In International Conference on Machine Learning, 2023.
  22. 22.Grosse, R. B., Ghahramani, Z., and Adams, R. P. Sandwiching the marginal likelihood using bidirectional monte carlo. arXiv preprint arXiv:1511.02543, 2015.
  23. 23.Grosse, R. B., Ancha, S., and Roy, D. Measuring the reliability of mcmc inference with bidirectional monte carlo. Advances in Neural Information Processing Systems, 2016.
  24. 24.Gu, S. S., Ghahramani, Z., and Turner, R. E. Neural adaptive sequential monte carlo. Advances in neural information processing systems, 28, 2015.
  25. 25.Guo, H., Tan, B., Liu, Z., Xing, E. P., and Hu, Z. Efficient (soft) q-learning for text generation with limited good data. arXiv preprint arXiv:2106.07704, 2021.
  26. 26.Gutmann, M. and Hyvarinen, A. Noise-contrastive estimation: A new estimation principle for unnormalized statistical models. In International conference on artificial intelligence and statistics, pp. 297–304. JMLR Workshop and Conference Proceedings, 2010.
  27. 27.Haarnoja, T., Zhou, A., Abbeel, P., and Levine, S. Soft actor-critic: Off-policy maximum entropy deep reinforcement learning with a stochastic actor. In International conference on machine learning. PMLR, 2018.
  28. 28.Heng, J., Bishop, A., Deligiannidis, G., and Doucet, A. Controlled sequential monte carlo. Annals of Statistics, 48(5), 2020.
  29. 29.Holtzman, A., Buys, J., Du, L., Forbes, M., and Choi, Y. The curious case of neural text degeneration. In International Conference on Learning Representations, 2019.
  30. 30.Hu, E. J., Jain, M., Elmoznino, E., Kaddar, Y., Lajoie, G., Bengio, Y., and Malkin, N. Amortizing intractable inference in large language models. arXiv preprint arXiv:2310.04363, 2023.
  31. 31.Khalifa, M., Elsahar, H., and Dymetman, M. A distributional approach to controlled text generation. arXiv preprint arXiv:2012.11635, 2020.
  32. 32.Khanov, M., Burapacheep, J., and Li, Y. ARGS: Alignment as reward-guided search. In The Twelfth International Conference on Learning Representations, 2024. URL https://openreview.net/forum?id=shgx0eqdw6.
  33. 33.Korbak, T., Elsahar, H., Kruszewski, G., and Dymetman, M. Controlling conditional language models without catastrophic forgetting. In International Conference on Machine Learning, pp. 11499–11528. PMLR, 2022a.
  34. 34.Korbak, T., Perez, E., and Buckley, C. L. Rl with kl penalties is better viewed as bayesian inference. arXiv preprint arXiv:2205.11275, 2022b.
  35. 35.Krause, B., Gotmare, A. D., McCann, B., Keskar, N. S., Joty, S., Socher, R., and Rajani, N. F. Gedi: Generative discriminator guided sequence generation. arXiv preprint arXiv:2009.06367, 2020.
  36. 36.Lawson, D., Tucker, G., Naesseth, C. A., Maddison, C., Adams, R. P., and Teh, Y. W. Twisted variational sequential monte carlo. In Third workshop on Bayesian Deep Learning (NeurIPS), 2018.
  37. 37.Lawson, D., Raventos, A., Warrington, A., and Linderman, S. Sixo: Smoothing inference with twisted objectives, 2022.
  38. 38.Levine, S. Reinforcement learning and control as probabilistic inference: Tutorial and review. arXiv preprint arXiv:1805.00909, 2018.
  39. 39.Lew, A. K., Zhi-Xuan, T., Grand, G., and Mansinghka, V. K. Sequential monte carlo steering of large language models using probabilistic programs. arXiv preprint arXiv:2306.03081, 2023.
  40. 40.Li, Y. Distilbert-base-uncased-finetuned-mnli-amazon-query-shopping, 2023. URL https://huggingface.co/LiYuan/amazon-review-sentiment-analysis.
  41. 41.Lioutas, V., Lavington, J. W., Sefas, J., Niedoba, M., Liu, Y., Zwartsenberg, B., Dabiri, S., Wood, F., and Scibior, A. Critic sequential monte carlo. In The Eleventh International Conference on Learning Representations, 2022.
  42. 42.Liu, A., Sap, M., Lu, X., Swayamdipta, S., Bhagavatula, C., Smith, N. A., and Choi, Y. Dexperts: Decoding-time controlled text generation with experts and anti-experts. In 59th Annual Meeting of the Association for Computational Linguistics and the 11th International Joint Conference on Natural Language Processing, 2021.
  43. 43.Liu, J., Cohen, A., Pasunuru, R., Choi, Y., Hajishirzi, H., and Celikyilmaz, A. Don’t throw away your value model! making ppo even better via value-guided monte-carlo tree search decoding. arXiv e-prints, pp. arXiv–2309, 2023.
  44. 44.Maddison, C. J., Lawson, J., Tucker, G., Heess, N., Norouzi, M., Mnih, A., Doucet, A., and Teh, Y. Filtering variational objectives. Advances in Neural Information Processing Systems, 30, 2017.
  45. 45.Mudgal, S., Lee, J., Ganapathy, H., Li, Y., Wang, T., Huang, Y., Chen, Z., Cheng, H.-T., Collins, M., Strohman, T., et al. Controlled decoding from language models. arXiv preprint arXiv:2310.17022, 2023.
  46. 46.Nachum, O., Norouzi, M., Xu, K., and Schuurmans, D. Bridging the gap between value and policy based reinforcement learning. Advances in neural information processing systems, 30, 2017.
  47. 47.Ouyang, L., Wu, J., Jiang, X., Almeida, D., Wainwright, C., Mishkin, P., Zhang, C., Agarwal, S., Slama, K., Ray, A., et al. Training language models to follow instructions with human feedback. Advances in Neural Information Processing Systems, 35:27730–27744, 2022.
  48. 48.Parshakova, T., Andreoli, J.-M., and Dymetman, M. Distributional reinforcement learning for energy-based sequential models. arXiv preprint arXiv:1912.08517, 2019.
  49. 49.Perez, E., Huang, S., Song, F., Cai, T., Ring, R., Aslanides, J., Glaese, A., McAleese, N., and Irving, G. Red teaming language models with language models. In Proceedings of the 2022 Conference on Empirical Methods in Natural Language Processing, pp. 3419–3448, 2022.
  50. 50.Phan, D., Hoffman, M. D., Douglas, S., Le, T. A., Parisi, A. T., Sountsov, P., Sutton, C., Vikram, S., Saurous, R. A., et al. Training chain-of-thought via latent-variable inference. In Thirty-seventh Conference on Neural Information Processing Systems, 2023.
  51. 51.Piche, A., Thomas, V., Ibrahim, C., Bengio, Y., and Pal, C. Probabilistic planning with sequential monte carlo methods. In International Conference on Learning Representations, 2018.
  52. 52.Qin, L., Welleck, S., Khashabi, D., and Choi, Y. Cold decoding: Energy-based constrained text generation with langevin dynamics. Advances in Neural Information Processing Systems, 35:9538–9551, 2022.
  53. 53.Radford, A., Wu, J., Child, R., Luan, D., Amodei, D., and Sutskever, I. Language models are unsupervised multitask learners. 2019. URL https://huggingface.co/gpt2-medium.
  54. 54.Rafailov, R., Sharma, A., Mitchell, E., Ermon, S., Manning, C. D., and Finn, C. Direct preference optimization: Your language model is secretly a reward model. arXiv preprint arXiv:2305.18290, 2023.
  55. 55.Scharth, M. and Kohn, R. Particle efficient importance sampling. Journal of Econometrics, 190(1):133–147, 2016.
  56. 56.Schulman, J., Wolski, F., Dhariwal, P., Radford, A., and Klimov, O. Proximal policy optimization algorithms. arXiv preprint arXiv:1707.06347, 2017.
  57. 57.Shih, A., Sadigh, D., and Ermon, S. Long horizon temperature scaling. arXiv preprint arXiv:2302.03686, 2023.
  58. 58.Snell, C. V., Kostrikov, I., Su, Y., Yang, S., and Levine, S. Offline rl for natural language generation with implicit language q learning. In The Eleventh International Conference on Learning Representations, 2022.
  59. 59.Sobolev, A. and Vetrov, D. P. Importance weighted hierarchical variational inference. Advances in Neural Information Processing Systems, 32, 2019.
  60. 60.Stiennon, N., Ouyang, L., Wu, J., Ziegler, D., Lowe, R., Voss, C., Radford, A., Amodei, D., and Christiano, P. F. Learning to summarize with human feedback. Advances in Neural Information Processing Systems, 33:3008–3021, 2020.
  61. 61.Vilnis, L., Zemlyanskiy, Y., Murray, P., Passos, A. T., and Sanghai, S. Arithmetic sampling: parallel diverse decoding for large language models. In International Conference on Machine Learning. PMLR, 2023.
  62. 62.Whiteley, N. and Lee, A. Twisted particle filters. The Annals of Statistics, 42(1):115–141, 2014.
  63. 63.Yang, K. and Klein, D. Fudge: Controlled text generation with future discriminators. In Proceedings of the 2021 Conference of the North American Chapter of the Association for Computational Linguistics: Human Language Technologies, pp. 3511–3535, 2021.
  64. 64.Zhang, H., Song, H., Li, S., Zhou, M., and Song, D. A survey of controllable text generation using transformer-based pre-trained language models. ACM Computing Surveys, 56(3):1–37, 2023.
  65. 65.Ziegler, D. M., Stiennon, N., Wu, J., Brown, T. B., Radford, A., Amodei, D., Christiano, P., and Irving, G. Fine-tuning language models from human preferences. arXiv preprint arXiv:1909.08593, 2019.
  66. 66.Zou, A., Wang, Z., Kolter, J. Z., and Fredrikson, M. Universal and transferable adversarial attacks on aligned language models. arXiv preprint arXiv:2307.15043, 2023.

Citation

MLA
Zhao, S., et al. “Probabilistic Inference in Language Models via Twisted Sequential Monte Carlo”. arXiv, 2024, http://arxiv.org/abs/2404.17546v1.
APA
Zhao, S., Brekelmans, R., Makhzani, A., & Grosse, R. (2024). Probabilistic Inference in Language Models via Twisted Sequential Monte Carlo. arXiv. http://arxiv.org/abs/2404.17546v1
Chicago
Zhao, S., R. Brekelmans, A. Makhzani, and R. Grosse. 2024. “Probabilistic Inference in Language Models via Twisted Sequential Monte Carlo”. arXiv. http://arxiv.org/abs/2404.17546v1.
Harvard
Zhao, S. et al. (2024) “Probabilistic Inference in Language Models via Twisted Sequential Monte Carlo”, arXiv [Preprint]. Available at: http://arxiv.org/abs/2404.17546v1.
Vancouver
1. Zhao S, Brekelmans R, Makhzani A, Grosse R (2024) Probabilistic Inference in Language Models via Twisted Sequential Monte Carlo. arXiv

BibTeX

@article{zhao2024probabilistic,
  title = {Probabilistic Inference in Language Models via Twisted Sequential Monte Carlo},
  author = {Zhao, Stephen and Brekelmans, Rob and Makhzani, Alireza and Grosse, Roger},
  year = {2024},
  journal = {arXiv},
  url = {http://arxiv.org/abs/2404.17546v1},
  eprint = {2404.17546}
}
Metadata:arXiv

Source Code

This paper has an official code repository available. Click below to access the source code.

View Repository

Access the Paper

This paper is available from its original source. Click below to access the PDF.

Open PDF
License: https://creativecommons.org/licenses/by/4.0/