File size: 1,724 Bytes
12496fc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 | """Tested loss primitives, not a claim of completed large-scale RL training."""
import torch
from torch.nn import functional as F
def masked_sft_loss(logits, labels, assistant_mask):
if labels.shape != assistant_mask.shape or logits.shape[:-1] != labels.shape:
raise ValueError("Incompatible SFT shapes")
if not assistant_mask.any():
raise ValueError("No assistant target tokens")
targets = labels.masked_fill(~assistant_mask.bool(), -100)
return F.cross_entropy(logits.flatten(0, -2), targets.flatten(), ignore_index=-100)
def dpo_loss(policy_chosen, policy_rejected, reference_chosen, reference_rejected, beta=0.1):
if beta <= 0:
raise ValueError("beta must be positive")
shapes = {x.shape for x in (policy_chosen, policy_rejected, reference_chosen, reference_rejected)}
if len(shapes) != 1:
raise ValueError("Log-probability shapes must match")
margin = (policy_chosen-policy_rejected) - (reference_chosen-reference_rejected)
return -F.logsigmoid(beta * margin).mean()
def group_advantages(rewards, eps=1e-6):
if rewards.ndim != 2 or rewards.shape[1] < 2 or not torch.isfinite(rewards).all():
raise ValueError("Expected finite batch x group rewards, group size >= 2")
return (rewards-rewards.mean(-1, keepdim=True)) / rewards.std(-1, keepdim=True, unbiased=False).clamp_min(eps)
def rejection_sample(candidates, verifier):
"""Return only independently verified candidates; never reward persuasive prose."""
accepted = []
for candidate in candidates:
try:
if verifier(candidate) is True:
accepted.append(candidate)
except Exception:
continue
return accepted
|