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()