File size: 4,248 Bytes
4397e12 | 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 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 | """GRPO reward shaping and credit assignment, as pure functions over per-episode numbers.
Settings follow the reference profile MiMo released with MiMo-V2.6 (recipes/arvo/REFERENCE_PENALTIES.json
in XiaomiMiMo/verl):
* reward: 1 correct and grounded, 0.5 correct but the answer never appeared in a tool result, 0 else;
* in-group length penalty on passing rollouts only, and only when more than half the group passes:
excess = max over (turns, tool-input tokens, generated tokens) of value / p30-of-passes - 1,
penalty = 0.2 * min(excess, 1) ** 1.5;
* advantage = shaped reward - group mean (no std normalization); groups whose shaped rewards are
all equal carry no signal and are dropped (dynamic sampling);
* segment-level penalty for bad tool-call turns (malformed, unknown tool, bad arguments, or a
repeated call): in positive episodes those tokens get no credit, in negative ones double blame,
and a batch-wide rescale of the clean tokens keeps each sign's total advantage mass unchanged.
"""
from __future__ import annotations
from dataclasses import dataclass
import numpy as np
@dataclass(frozen=True)
class LengthPenalty:
max_penalty: float = 0.2
threshold: float = 0.0 # tolerated excess before any penalty
saturate: float = 1.0 # excess at which the penalty is maximal (2x the anchor)
exponent: float = 1.5
pass_threshold: float = 0.5
anchor_quantile: float = 0.3
min_pass_rate: float = 0.5 # apply only when pass fraction is strictly above this
def base_reward(correct: bool, grounded: bool, ungrounded: float = 0.5, invented: bool = False,
invented_penalty: float = 0.0) -> float:
"""1 = correct and read from a tool result, `ungrounded` = correct but never seen, 0 = wrong,
-invented_penalty = wrong and seen nowhere (made up). So within a group, admitting NOT_FOUND
beats inventing an answer, while a correct answer still beats both."""
if correct:
return 1.0 if grounded else ungrounded
return -invented_penalty if invented else 0.0
def length_deltas(rewards: list[float], signals: list[dict], cfg: LengthPenalty = LengthPenalty()) -> list[float]:
"""Non-positive reward deltas for one group. signals: per episode {metric: value}."""
n = len(rewards)
out = [0.0] * n
passed = [i for i in range(n) if rewards[i] >= cfg.pass_threshold]
if not passed or len(passed) / n <= cfg.min_pass_rate:
return out
metrics = signals[0].keys()
anchor = {m: float(np.percentile([signals[i][m] for i in passed], cfg.anchor_quantile * 100)) for m in metrics}
for i in passed:
ex = [max(0.0, signals[i][m] / anchor[m] - 1.0) for m in metrics if anchor[m] > 0]
e = max(ex, default=0.0)
if e <= cfg.threshold:
continue
t = min((e - cfg.threshold) / (cfg.saturate - cfg.threshold), 1.0)
out[i] = -cfg.max_penalty * t ** cfg.exponent
return out
def group_advantages(rewards: list[float]) -> list[float]:
mu = sum(rewards) / len(rewards)
return [r - mu for r in rewards]
def has_signal(rewards: list[float], eps: float = 1e-9) -> bool:
return max(rewards) - min(rewards) > eps
def signed_rebalance(adv: list[float], n_flag: list[int], n_clean: list[int], kappa: float = 2.0,
min_scale: float = 0.5, max_scale: float = 2.0) -> list[tuple[float, float]]:
"""Per episode (weight on clean tokens, weight on flagged tokens); token advantage = A * weight.
Batch-wide: positive episodes drop credit on flagged tokens and scale clean ones by alpha;
negative episodes put kappa x blame on flagged tokens and scale clean ones by beta, so that
(unless clipped) the total positive and negative advantage mass is conserved."""
hp = sum(a * f for a, f in zip(adv, n_flag) if a > 0)
cp = sum(a * c for a, c in zip(adv, n_clean) if a > 0)
hn = sum(-a * f for a, f in zip(adv, n_flag) if a < 0)
cn = sum(-a * c for a, c in zip(adv, n_clean) if a < 0)
alpha = min(max_scale, 1 + hp / cp) if cp > 0 else 1.0
beta = max(min_scale, 1 - (kappa - 1) * hn / cn) if cn > 0 else 1.0
return [(alpha, 0.0) if a > 0 else (beta, kappa) if a < 0 else (0.0, 0.0) for a in adv]
|