Tokenize (prompt, chosen, rejected)
Forward pass: policy on chosen+rejected
Forward pass: frozen reference (or cached)
Sum token log-probs over response spans
Compute DPO loss
python
import torch
import torch.nn.functional as F
def response_logprobs(logits, labels, prompt_len):
"""
logits: [B, T, V] from the model
labels: [B, T] input token ids (shifted target = labels[:, 1:])
prompt_len: [B] length of the prompt portion to mask out
Returns: [B] summed log-prob of the response tokens only.
"""
logp = F.log_softmax(logits[:, :-1, :], dim=-1) # predict token t from t-1
targets = labels[:, 1:]
token_logp = torch.gather(logp, 2, targets.unsqueeze(-1)).squeeze(-1) # [B, T-1]
T = token_logp.shape[1]
positions = torch.arange(T, device=logits.device).unsqueeze(0)
response_mask = positions >= (prompt_len.unsqueeze(1) - 1) # mask out prompt tokens
return (token_logp * response_mask).sum(dim=1)