NEXORA / nexora /posttraining.py
devildasdf's picture
Release validated NEXORA research prototype, tiny weights and evidence
12496fc verified
Raw History Blame Contribute Delete
1.72 kB
"""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