Skip to content
AILinkDeepTech
Go back
Deep Learning Medium

Direct Preference Optimization (DPO) Implementation in PyTorch

Abstract

A PyTorch implementation of Direct Preference Optimization (DPO): a small LSTM language model, a preference dataset, and a DPO trainer that optimizes a policy against a frozen reference via the sigmoid logistic DPO loss on chosen/rejected pairs.

Direct Preference Optimization (DPO) Implementation in PyTorch

This implementation builds Direct Preference Optimization (DPO) from scratch. It includes a small LSTM language model, a synthetic preference dataset with hash-based tokenization, and a DPOTrainer that optimizes a policy model against a frozen reference using the sigmoid logistic DPO loss on chosen/rejected pairs.

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader
from tqdm import tqdm

class SimpleLanguageModel(nn.Module):
    """
    A very simple language model for testing DPO implementation.
    This is much smaller than GPT-2 and suitable for local testing.
    """
    def __init__(self, vocab_size=1000, hidden_size=32, num_layers=2):
        super(SimpleLanguageModel, self).__init__()
        self.embedding = nn.Embedding(vocab_size, hidden_size)
        self.lstm = nn.LSTM(
            input_size=hidden_size,
            hidden_size=hidden_size,
            num_layers=num_layers,
            batch_first=True
        )
        self.fc = nn.Linear(hidden_size, vocab_size)

    def forward(self, input_ids, attention_mask=None):

        embedded = self.embedding(input_ids)

        if attention_mask is not None:
            lengths = attention_mask.sum(dim=1).cpu()
            packed_embedded = nn.utils.rnn.pack_padded_sequence(
                embedded, lengths, batch_first=True, enforce_sorted=False
            )
            packed_output, _ = self.lstm(packed_embedded)
            lstm_out, _ = nn.utils.rnn.pad_packed_sequence(packed_output, batch_first=True)
        else:
            lstm_out, _ = self.lstm(embedded)

        logits = self.fc(lstm_out)
        
        return logits
    
class DPOTrainer:
    """
    Implementation of Direct Preference Optimization (DPO).
    DPO trains a policy model using a reference model and human preference data.
    """
    def __init__(
        self,
        policy_model,
        ref_model,
        beta=0.1,
        device="cuda" if torch.cuda.is_available() else "cpu",
        log_interval=5
    ):
        """     
        Args:
            policy_model: Model to be trained 
            ref_model: Reference model (kept frozen)
            beta: Temperature parameter for the DPO loss
            log_interval: How often to log training statistics
        """
        self.policy_model = policy_model.to(device)
        self.ref_model = ref_model.to(device)
        self.beta = beta
        self.device = device
        self.log_interval = log_interval

        # Freeze the reference model parameters
        for param in self.ref_model.parameters():
            param.requires_grad = False

        self.loss_history = []
        self.accuracy_history = []

    def _compute_response_logprobs(self, model, response_ids, response_mask):
        """
        Compute log probabilities of responses.
        """
        logits = model(response_ids, response_mask)  
        
        log_probs = F.log_softmax(logits, dim=2)  

        batch_size, seq_len = response_ids.size()
        
        total_log_prob = torch.zeros(batch_size, device=response_ids.device)
        token_counts = torch.zeros(batch_size, device=response_ids.device)

        for pos in range(seq_len - 1):
            # Check if current and next positions are valid (not padding)
            current_valid = response_mask[:, pos].bool()
            next_valid = response_mask[:, pos + 1].bool()
            valid = current_valid & next_valid
            
            next_token_ids = response_ids[:, pos + 1]

            # get log probability of the next token
            for i in range(batch_size):
                if valid[i]:
                    next_token = next_token_ids[i]
                    token_log_prob = log_probs[i, pos, next_token]
                    total_log_prob[i] += token_log_prob
                    token_counts[i] += 1

        token_counts = token_counts.clamp(min=1)
        
        avg_log_prob = total_log_prob / token_counts
        
        return avg_log_prob
    
    def train_step(self, batch_chosen_responses, batch_chosen_masks, 
                  batch_rejected_responses, batch_rejected_masks):
        # Compute log probabilities from policy model
        chosen_policy_logps = self._compute_response_logprobs(
            self.policy_model,
            batch_chosen_responses, 
            batch_chosen_masks
        )
        
        rejected_policy_logps = self._compute_response_logprobs(
            self.policy_model,
            batch_rejected_responses, 
            batch_rejected_masks
        )

        # Compute log probabilities from reference model
        with torch.no_grad():
            chosen_ref_logps = self._compute_response_logprobs(
                self.ref_model,
                batch_chosen_responses, 
                batch_chosen_masks
            )
            
            rejected_ref_logps = self._compute_response_logprobs(
                self.ref_model,
                batch_rejected_responses, 
                batch_rejected_masks
            )

        chosen_rewards = chosen_policy_logps - chosen_ref_logps
        rejected_rewards = rejected_policy_logps - rejected_ref_logps
        
        logits = self.beta * (chosen_rewards - rejected_rewards)

        # Target is always 1 (chosen should be preferred)
        targets = torch.ones_like(logits)
        
        loss = F.binary_cross_entropy_with_logits(logits, targets)
        
        accuracy = (logits > 0).float().mean()
        
        return loss, accuracy
    
    def train(self, dataloader, optimizer, num_epochs=1):
        """
        Train the policy model using DPO.
        """
        self.policy_model.train()
        self.ref_model.eval()

        for epoch in range(num_epochs):
            epoch_loss = 0
            epoch_acc = 0
            num_batches = 0
            
            progress_bar = tqdm(dataloader, desc=f"Epoch {epoch+1}/{num_epochs}")

            for batch in progress_bar:
                batch_chosen_responses = batch["chosen_responses"].to(self.device)
                batch_chosen_masks = batch["chosen_masks"].to(self.device)
                
                batch_rejected_responses = batch["rejected_responses"].to(self.device)
                batch_rejected_masks = batch["rejected_masks"].to(self.device)
                
                optimizer.zero_grad()

                # Forward pass and compute loss
                loss, accuracy = self.train_step(
                    batch_chosen_responses, batch_chosen_masks,
                    batch_rejected_responses, batch_rejected_masks
                )
                
                loss.backward()
                optimizer.step()

                # Update statistics
                epoch_loss += loss.item()
                epoch_acc += accuracy.item()
                num_batches += 1

                # Update progress bar
                if num_batches % self.log_interval == 0 or num_batches == len(dataloader):
                    progress_bar.set_postfix({
                        "loss": epoch_loss / num_batches,
                        "accuracy": epoch_acc / num_batches
                    })

            avg_loss = epoch_loss / num_batches
            avg_acc = epoch_acc / num_batches
            print(f"Epoch {epoch+1}/{num_epochs}: Avg Loss = {avg_loss:.4f}, Avg Accuracy = {avg_acc:.4f}")
            
            # Save history
            self.loss_history.append(avg_loss)
            self.accuracy_history.append(avg_acc)

        return {"loss": self.loss_history, "accuracy": self.accuracy_history}
    
class PreferenceDataset(Dataset):
    """
    Dataset for preference data used in DPO training.
    
    Each item contains:
    - A chosen (preferred) response
    - A rejected response
    """
    def __init__(self, chosen_texts, rejected_texts, vocab_size=1000, max_length=30):
        """
        Args:
            chosen_texts: List of chosen/preferred responses
            rejected_texts: List of rejected responses
            vocab_size: Size of vocabulary for simple tokenization
            max_length: Maximum sequence length
        """
        self.chosen_texts = chosen_texts
        self.rejected_texts = rejected_texts
        self.vocab_size = vocab_size
        self.max_length = max_length
        
        self.tokenize = lambda text: self._simple_tokenize(text, max_length)

    def _simple_tokenize(self, text, max_length):
        """
        Very simple tokenization by hashing words to indices.
        
        Args:
            text: Text to tokenize
            max_length: Maximum sequence length
        """
        words = text.lower().replace('.', ' .').replace(',', ' ,').replace('!', ' !').replace('?', ' ?').split()
        
        # Hash words to token IDs (reserve 0 for padding, 1 for unknown)
        tokens = [hash(word) % (self.vocab_size - 2) + 2 for word in words]
        
        if len(tokens) > max_length:
            tokens = tokens[:max_length]
            
        # Create mask (1 for tokens, 0 for padding)
        mask = [1] * len(tokens)
        
        # Pad to max length
        padding_needed = max_length - len(tokens)
        if padding_needed > 0:
            tokens = tokens + [0] * padding_needed
            mask = mask + [0] * padding_needed
            
        return tokens, mask
    
    def __len__(self):
        return len(self.chosen_texts)
    
    def __getitem__(self, idx):

        chosen = self.chosen_texts[idx]
        rejected = self.rejected_texts[idx]
        
        # Tokenize
        chosen_tokens, chosen_mask = self.tokenize(chosen)
        rejected_tokens, rejected_mask = self.tokenize(rejected)
        
        chosen_tokens = torch.tensor(chosen_tokens)
        chosen_mask = torch.tensor(chosen_mask)
        rejected_tokens = torch.tensor(rejected_tokens)
        rejected_mask = torch.tensor(rejected_mask)
        
        return {
            "chosen_responses": chosen_tokens,
            "chosen_masks": chosen_mask,
            "rejected_responses": rejected_tokens,
            "rejected_masks": rejected_mask
        }
    
def run_simple_test():
    print("Starting DPO test case with simple models...")
    
    device = "cuda" if torch.cuda.is_available() else "cpu"
    print(f"Using device: {device}")
    
    print("Creating small language models...")
    vocab_size = 1000
    hidden_size = 32
    
    policy_model = SimpleLanguageModel(vocab_size=vocab_size, hidden_size=hidden_size)
    ref_model = SimpleLanguageModel(vocab_size=vocab_size, hidden_size=hidden_size)

    # Copy weights from ref_model to policy_model 
    ref_model_state = ref_model.state_dict()
    policy_model.load_state_dict(ref_model_state)

    print("Creating synthetic preference dataset...")
    
    chosen_responses = [
        "Once upon a time there was a magnificent dragon with emerald scales.",
        "Quantum computing leverages quantum mechanical phenomena like superposition.",
        "The best way to learn programming is to work on projects that interest you.",
        "The vast blue expanse stretches beyond sight waves dancing beneath light.",
    ] * 5
    
    rejected_responses = [
        "Dragon story big lizard breathes fire.",
        "Quantum computers use qubits instead of bits.",
        "Just read some books and learn programming.",
        "Ocean is wet. Ocean is blue. Fish swim there too.",
    ] * 5
    
    dataset = PreferenceDataset(chosen_responses, rejected_responses)
    dataloader = DataLoader(dataset, batch_size=4, shuffle=True)

    print("Initializing DPO trainer...")
    dpo_trainer = DPOTrainer(
        policy_model=policy_model,
        ref_model=ref_model,
        beta=0.1,
        device=device
    )
    
    optimizer = optim.AdamW(policy_model.parameters(), lr=5e-4)

    print("Training model with DPO...")
    history = dpo_trainer.train(dataloader, optimizer, num_epochs=5)
    
    print("\nTraining complete!")
    print(f"Final loss: {history['loss'][-1]:.4f}")
    print(f"Final accuracy: {history['accuracy'][-1]:.4f}")

    # Test the model on a sample prompt
    test_chosen = "The future of AI will likely involve more transparent interpretable models."
    test_rejected = "AI will take over everything soon."
    
    chosen_tokens, chosen_mask = dataset.tokenize(test_chosen)
    rejected_tokens, rejected_mask = dataset.tokenize(test_rejected)
    
    chosen_tokens = torch.tensor([chosen_tokens]).to(device)
    chosen_mask = torch.tensor([chosen_mask]).to(device)
    rejected_tokens = torch.tensor([rejected_tokens]).to(device)
    rejected_mask = torch.tensor([rejected_mask]).to(device)

    # Get policy and reference scores
    with torch.no_grad():
        policy_chosen_logp = dpo_trainer._compute_response_logprobs(
            policy_model,
            chosen_tokens,
            chosen_mask
        )
        
        policy_rejected_logp = dpo_trainer._compute_response_logprobs(
            policy_model,
            rejected_tokens,
            rejected_mask
        )
        
        ref_chosen_logp = dpo_trainer._compute_response_logprobs(
            ref_model,
            chosen_tokens,
            chosen_mask
        )
        
        ref_rejected_logp = dpo_trainer._compute_response_logprobs(
            ref_model,
            rejected_tokens,
            rejected_mask
        )

    print("\nTesting model preferences:")
    print(f"Policy model preference: {'Chosen' if policy_chosen_logp > policy_rejected_logp else 'Rejected'}")
    print(f"  - Chosen log prob: {policy_chosen_logp.item():.4f}")
    print(f"  - Rejected log prob: {policy_rejected_logp.item():.4f}")

    chosen_reward = (policy_chosen_logp - ref_chosen_logp).item()
    rejected_reward = (policy_rejected_logp - ref_rejected_logp).item()
    
    print("\nImplicit rewards:")
    print(f"  - Chosen reward: {chosen_reward:.4f}")
    print(f"  - Rejected reward: {rejected_reward:.4f}")

    if chosen_reward > rejected_reward:
        print("\nSUCCESS! Model learned to assign higher rewards to chosen responses.")
    else:
        print("\nTraining may need more epochs or hyperparameter tuning.")
    
    print("\nTest complete!")

if __name__ == "__main__":
    run_simple_test()


Cite this Explanation

@article{ailinkdeeptech2025dpoalgo,
  title={Direct Preference Optimization (DPO) Implementation in PyTorch},
  author={AILinkDeepTech},
  journal={AILinkDeepTech Algorithm Explanations},
  year={2025},
  url={https://ailinkdeeptech.com/research/dpo_algo}
}

Related Explanations