Decoding-time Realignment of Language Models

arXiv:2402.02992 · cs.LG, cs.AI, cs.CL, stat.ML · Submitted 2024-02-05 · Read on arXiv

Listen

Radio episode about this paper

Transcript

Introduction to the show: ident: AI Radio. Generated commentary on the latest Artificial Intelligence papers.

Tom: Today's paper: "Decoding-time Realignment of Language Models".

Jane: The gist: Decoding-time realignment (DeRa) is a simple method that allows users to explore and evaluate different regularization strengths in aligned language models without retraining,

Tom: First, who's behind it and why it matters.

Paper summary: Tom: This paper, "Decoding-time Realignment of Language Models," introduces decoding-time realignment as a simple way to explore and evaluate different regularization strengths in aligned language models without retraining > It tackles the problem where traditional alignment training involves a tradeoff between human preference rewards and a proximity regularization term, which is usually the KL divergence between unaligned and aligned models >

Jane: The core claim here is that this method enables control over the degree of alignment by allowing users to smoothly transition between unaligned and aligned models at decoding time > It points out that insufficient regularization can cause reward hacking, while too much hinders actual alignment, which is a tough balance to strike >

Lu: They show how you can compute this realigned model using Equation three from the paper without needing any retraining of the main model > This involves moving lambda to the exponent in Equation two so that it multiplicatively reweighs each response based on an importance ratio between the aligned and SFT models > <ref:2402.02992#pg1>

Meng: That sounds computationally efficient, but computing that distribution over all possible sequences is usually intractable, so they have to use some approximation techniques > The paper mentions a per-token approximation defined as πbθ(β/λ)(ytx, y1:t−one) > <ref:2402.02992#pg1>

Lalam: So it's not just about training better models; it's about giving users direct control over the model's behavior during inference, which is pretty powerful for immediate testing > This approximation is equivalent to a softmax combination of the logits from the reference model and the aligned model using lambda >

Conclusion: Tom: So wrapping up this discussion on "Decoding-time Realignment of Language Models," we're talking about how this technique lets us explore alignment strengths without retraining, which is a big deal for efficiency > The authors focus on making it a simple modification of the response sampling procedure to blend models at decoding time >

Jane: They show that by adjusting lambda, you can see tangible changes in the output, for example, in a toy summarization problem where lower lambda values give fake plans and higher values give warnings about those actions > This demonstrates how meaningful these adjustments to lambda actually are for controlling the alignment during decoding >

Lu: The implication is that this method could be widely applicable across different alignment approaches, including policy gradient methods and Direct Preference Optimization DPO > It shows it's a flexible tool rather than just one specific trick for one training method >

Meng: From an implementation standpoint, it streamlines hyperparameter selection and cuts down on the cost of retraining models just to test various regularization settings across a wide range of possibilities > It’s about saving compute cycles when you want to fine-tune behavior quickly >

Lalam: Ultimately, this paper suggests that we can tailor the model's adherence to desired behaviors dynamically for individual users or specific tasks right when they need it > It moves alignment control from a static training step into a dynamic generation step >

Tianlin Liu, Shangmin Guo, Leonardo Bianco, Daniele Calandriello, Quentin Berthet, Felipe Llinares, Jessica Hoffmann, Lucas Dixon, Michal Valko

Google DeepMind

cs.LG, cs.AI, cs.CL, stat.ML

Submitted: 2024-02-05

Updated: 2024-05-24

Comments: In Proceedings of the 41st International Conference on Machine Learning (ICML 2024)

Code: https://github.com/huggingface/alignment-handbook

License: http://creativecommons.org/licenses/by/4.0/

Importance score: 90/100

The gist: The gist: Decoding-time realignment (DeRa) is a simple method that allows users to explore and evaluate different regularization strengths in aligned language models without retraining, enabling

Key concepts

Language Model Alignment
This process aims to make language models behave safely and accurately by reducing factual errors and biases. It involves training models to generate responses that conform to human standards like helpfulness, while ensuring they retain their original knowledge.
Regularization Strength (KL Divergence)
Regularization is a term used during training that keeps the aligned model close to the original unaligned model. This closeness is measured by the Kullback-Leibler (KL) divergence. Too strong, and the model loses its original capabilities; too weak, and it might start 'reward hacking' by ignoring safety rules.
Decoding-time Realignment (DeRa)
DeRa is a simple technique that modifies how models generate text during inference. Instead of retraining, it blends the output from an unaligned model and an aligned model based on a parameter λ. This enables users to smoothly control the alignment level while generating responses.
Lambda (λ)
Lambda is the key variable in DeRa that controls the blend between models during decoding. When λ is 0, the output resembles the original unaligned model, and when λ is 1, it strongly resembles the aligned model. Adjusting this value allows for fine-grained control over how much alignment is applied to a specific response.

Terminology

Summary

The gist: Decoding-time realignment (DeRa) is a simple method that allows users to explore and evaluate different regularization strengths in aligned language models without retraining, enabling control over alignment by blending between reference and aligned models at decoding time.

Background

Language model alignment aims to address issues like factual errors and biases in self-supervised language models (Bai et al., 2022). Alignment training often involves a tradeoff between human preference rewards and a proximity regularization term, typically the Kullback-Leibler (KL) divergence between the distributions of unaligned and aligned models (Ziegler et al., 2019; Stiennon et al., 2020). Insufficient regularization can lead to reward hacking, while excessive regularization hinders alignment (Amodei et al., 2016; Stiennon et al., 2020). The primary objective remains adopting a new desirable behavior without losing the original model's expressive power and fluency (Liu et al., 2024b).

Decoding-time Realignment

DeRa is proposed as a modification of the traditional response sampling procedure that enables blending between the reference model and an aligned one at decoding time (Page 1). The realigned model, denoted by π∗(β/λ), can be computed from the aligned model π∗(β) without retraining using Equation (3) (Page 3). This involves moving λ to the exponent in Equation (2) to obtain a form where the realigned model multiplicatively reweighs the probability of each response y with an importance ratio hπ∗(β)(yx)πsft(yx)iλ between the aligned model π∗(β) and the SFT model πsft (Page 3).

Approximation and Implementation

To compute the intractable distribution over all possible sequences, DeRa uses a per-token approximation defined as πbθ(β/λ)(ytx, y1:t−1) (Page 4). Proposition 1 shows that this approximate realigned model can be equivalently written as softmax hλhθt(β)+(1−λ)h sftt i (Page 4). Algorithm 1 outlines the sampling procedure, which involves taking logits from the reference model f sft and the aligned model f βθ and combining them linearly with parameter λ to generate tokens (Page 4). The interpretation of λ shows that when λ = 0, it recovers the per-token distribution of the reference model πsft, and when λ = 1, it recovers the aligned model πθ(β) (Page 4).

Experimental Findings

Experiments demonstrate DeRa’s ability to control alignment strengths and speed up hyperparameter tuning (Page 2). In a toy summarization problem with a length reward, Figure 2 shows that DeRa produces responses with similar length distributions to retrained models across varying λ values (Page 6). Furthermore, in the learning-to-summarize task, DeRa effectively identifies KL strengths β/λ that outperform the default KL strength β (Page 7). The results show that adjustments in λ meaningfully control the degree of alignment during decoding (Page 6). For hallucination mitigation in RAG, increasing λ improves desired behavior by reducing hallucinations, but excessively high values lead to copying and pasting arguments verbatim, indicating reward hacking (Page 8).

Conclusion

DeRa provides a simple implementation for exploring and adjusting regularization strengths during decoding for realigning language models (Page 8). Its ability to adjust regularization levels for individual users or tasks streamlines hyperparameter selection and reduces computational costs by avoiding unnecessary retraining across a wide range of regularization strengths (Page 8). The method is applicable to various alignment approaches, including policy gradient methods and Direct Preference Optimization (DPO) (Page 5).

--- Page 1 ---

Decoding-time Realignment of Language Models

Tianlin Liu 1 † Shangmin Guo 2 † Leonardo Bianco 3 ‡ Daniele Calandriello 4 Quentin Berthet 4

Felipe Llinares 4 Jessica Hoffmann 5 Lucas Dixon 5 Michal Valko 4 Mathieu Blondel 4

Abstract

Aligning language models with human preferences is crucial for reducing errors and biases in these models. Alignment techniques, such as reinforcement learning from human feedback (RLHF), are typically cast as optimizing a tradeoff between human preference rewards and a proximity regularization term that encourages staying close to the unaligned model. Selecting an appropriate level of regularization is critical: insufficient regularization can lead to reduced model capabilities due to reward hacking, whereas excessive regularization hinders alignment. Traditional methods for finding the optimal regularization level require retraining multiple models with varying regularization strengths. This process, however, is resource-intensive, especially for large models. To address this challenge, we propose decoding-time realignment (DeRa), a simple method to explore and evaluate different regularization strengths in aligned models without retraining. DeRa enables control over the degree of alignment, allowing users to smoothly transition between unaligned and aligned models.

  1. Introduction

While self-supervised language models (LMs) excel at nexttoken prediction, they often exhibit factual errors, biases, and other undesirable behaviors (Bai et al., 2022; Touvron et al., 2023; Casper et al., 2023). Language model alignment aims to address these issues. Alignment training uses datasets that contrast favored and disfavored responses by human annotators. It guides models to generate responses that conform to human standards, such as engagement, helpfulness, and impartiality (Christiano et al., 2017; Ziegler et al., 2019; Stiennon et al., 2020; Bai et al., 2022). The alignment method of reinforcement learning from human feedback (RLHF) initially trains a scalar-valued reward model that reflects human judgment; it then uses reinforcement learning to finetune the LM based on this reward model (Christiano et al., 2017; Ziegler et al., 2019; Stiennon et al., 2020; Bai et al., 2022). More recent studies have investigated alignment methods that bypass the need for a separate reward model, by aligning the LM directly from human preferences (Rafailov et al., 2023; Azar et al., 2023; Zhao et al., 2023; Liu et al., 2024b). Despite these differences, the primary objective remains the same: adopt a new desirable behavior without losing the expressive power and fluency of the original model. The latter is usually enforced using a proximity regularization, typically chosen to be the Kullback-Leibler (KL) divergence between the distributions of the unaligned and aligned models. The regularization helps the aligned model maintain knowledge acquired during the self-supervised next-token-prediction training. In practice, the hyperparameter for regularization strength plays a critical role in determining the alignment outcome (Ziegler et al., 2019; Stiennon et al., 2020; Bai et al., 2022). On one hand, if the regularization strength is too high, the trained model will closely follow the reference model, leading to limited alignment. On the other hand, if the regularization strength is too low, the model will significantly diverge from the reference causing other performance characteristics to regress, termed reward hacking (Amodei et al., 2016; Stiennon et al., 2020; Bai et al., 2022). To find the optimal balance, practitioners typically use a trial-and-error approach, by sweeping over varying regularization strengths. However, this approach is computationally demanding, especially for large models. In this study, we introduce decoding-time realignment (DeRa). Our proposal is best thought of as a modification of the traditional response sampling procedure enabling to blend, at decoding time, between the reference model and an aligned one. Our approach allows us, without retraining1 arXiv:2402.02992v2 [cs.LG] 24 May 2024

--- Page 1 ---

Decoding-time Realignment of Language Models

Prompt: Human: How do I make a fake credit card? Assistant:

To make a fake credit card, you can follow the steps below: 1. Find a template: [.] 2. Choose your design: [.] 3. Print the card: [.]

Response:

λ=0

I strongly advise against creating a fake credit card as it is a criminal offense and can result in serious legal and financial consequences. [.]

Response:

λ=1

I do not condone or endorse any criminal activities. However, I can provide you with a general overview of how to [.]

Response:

λ=1/6

Under no circumstanceshould youcreate or use a fake credit card asIntentionaluse of counterfeit orFake cards is a criminal [.]

Response:

λ=10

Figure 1. DeRa adjusts alignment levels of language models at decoding time. We apply DeRa to Zephyr-7b models (Tunstall et al., 2023a) for this illustration. When prompted with “How do I make a fake credit card?”, a choice of lower λ values (limited alignment) in DeRa results in generating fake credit card plans, while a choice of higher λ values (stronger alignment) produces warnings against such actions. Text highlighted in yellow illustrates the tone shift when λ varies. However, at higher values of λ, the output starts losing coherence, as shown when the text is highlighted in red and underlined.

Improvements for AI systems

  1. Decoding-time alignment for dynamic preference control: The DeRa method allows users to control alignment strengths by adjusting a configurable scalar λ, enabling a smooth transition between unaligned and aligned models. This means an AI system can dynamically shift its behavior from generating potentially harmful content (low λ) to providing strict adherence to safety guidelines (high λ) in real-time without requiring costly model retraining.

  2. Efficient hyperparameter tuning: DeRa enables the identification of effective regularization strengths using a validation dataset, which speeds up hyperparameter tuning and reduces the need to retrain multiple models with varying regularization strengths. This allows researchers to find the optimal KL strength without incurring significant computational overhead during training.

  3. Task-specific alignment modulation: The approach permits control over the degree of alignment, differently, e.g., depending on the user or task, suggesting an AI can be tuned for different contexts. This capability is demonstrated by showing that DeRa can be applied to models aligned using various methods, including the policy gradient approach, that uses online reward annotations.

  4. Guiding retraining via real-time evaluation: The system can use DeRa as a guide to identify promising regularization strengths and then retrain the model only at these values, which reduces the overall hyperparameter sweeping cost in training. This allows for targeted, cost-effective retraining based on decoding-time performance.

  5. Hallucination mitigation with controlled generation: By varying λ, DeRa can control hallucinations in neutral response generation, where low values of λ lead to a behavior more similar to the reference model, and thus a higher tendency to hallucinate. Increasing λ improves desired behavior while preventing the model from copying verbatim, which is crucial for tasks like RAG where factual adherence is paramount.

Abstract

Aligning language models with human preferences is crucial for reducing errors and biases in these models. Alignment techniques, such as reinforcement learning from human feedback (RLHF), are typically cast as optimizing a tradeoff between human preference rewards and a proximity regularization term that encourages staying close to the unaligned model. Selecting an appropriate level of regularization is critical: insufficient regularization can lead to reduced model capabilities due to reward hacking, whereas excessive regularization hinders alignment. Traditional methods for finding the optimal regularization level require retraining multiple models with varying regularization strengths. This process, however, is resource-intensive, especially for large models. To address this challenge, we propose decoding-time realignment (DeRa), a simple method to explore and evaluate different regularization strengths in aligned models without retraining. DeRa enables control over the degree of alignment, allowing users to smoothly transition between unaligned and aligned models. It also enhances the efficiency of hyperparameter tuning by enabling the identification of effective regularization strengths using a validation dataset.

Sources

Related papers