"""DPO 损失与 GRPO 组内相对优势。"""
import torch
import torch.nn.functional as F


def dpo_loss(pc, pr, rc, rr, beta=0.1):
    logits=beta*((pc-pr)-(rc-rr))
    return -F.logsigmoid(logits).mean()


def group_advantage(rewards, eps=1e-6):
    mean=rewards.mean(1,keepdim=True)
    std=rewards.std(1,keepdim=True,unbiased=False)
    return (rewards-mean)/(std+eps)


if __name__ == "__main__":
    loss=dpo_loss(torch.tensor([-2.]),torch.tensor([-4.]),torch.tensor([-3.]),torch.tensor([-3.5]))
    adv=group_advantage(torch.tensor([[1.,2.,3.],[4.,4.,4.]]))
    assert loss.ndim==0 and torch.allclose(adv.sum(1),torch.zeros(2),atol=1e-5)
    print("loss",loss.item(),"adv",adv)

