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