Download nexora/posttraining.py from devildasdf/NEXORA: direct link, hf CLI and curl.
- Browser
- Download file 1.72 kB
-
https://huggingface.co/devildasdf/NEXORA/resolve/main/nexora/posttraining.py
- Command line
-
hf download hf://devildasdf/NEXORA/nexora/posttraining.py
-
curl -L -o posttraining.py https://huggingface.co/devildasdf/NEXORA/resolve/main/nexora/posttraining.py
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 | |