Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
reward.py247 linesDownload Raw Back to altered_minds
1"""reward.py — MMLUFormatReward (ADR-013, framework-side, generic).2 3A structured-answer reward for RL on MMLU-style multiple-choice tasks. It scores4ONLY the final answer letter + format validity — never the rationale's style or5content. This is deliberate: the north-star use case (ADR-013) drives RL on a6*personality-altered* model, and rewarding "persuasive" rationale would reward7the very cognitive-distortion signature we are trying to measure rather than8distort the reward toward it.9 10Scoring per completion:11  +1.0   final answer parses and equals the gold letter12   0.0   final answer parses but is wrong13  -0.2   no parseable final-answer marker (unparseable)14  -0.1   multiple DISTINCT final-answer markers present (format hacking)15  -len_penalty   small penalty past a rationale character cap16 17Parsing accepts (case-insensitive, last match wins for the canonical letter):18  - ``Answer: X``           (X in A-D)19  - JSON ``{"answer": "X"}``20 21Exploit detection: ``MMLUFormatReward`` keeps a running count of chosen letters22(``option_distribution``) so an "always C" / option-prior exploit is detectable23by inspecting that distribution after a run. A companion ``randomize_options``24helper shuffles option order with an original->shuffled label remap so the25training data itself can be de-biased.26"""27from __future__ import annotations28 29import json30import re31from collections import Counter32from dataclasses import dataclass, field33from typing import Any34 35__all__ = ["MMLUFormatReward", "randomize_options", "parse_final_answer"]36 37_VALID_LETTERS = ("A", "B", "C", "D")38 39# ``Answer: X`` — tolerant of whitespace, optional markdown bold/asterisks.40_ANSWER_RE = re.compile(r"answer\s*[:\-]\s*\*{0,2}([A-D])\b", re.IGNORECASE)41# JSON ``{"answer": "X"}`` — extract the value of an "answer" key.42_JSON_ANSWER_RE = re.compile(43    r'["\']answer["\']\s*:\s*["\']([A-D])["\']', re.IGNORECASE44)45# Hedge AFTER an answer marker: ``Answer: A or B`` / ``Answer: A/B`` /46# ``Answer: A, B`` — a single marker that names a SECOND distinct option is a47# format hedge and must be treated as multiple-answers, not full credit for the48# first letter (final-verify 2026-05-29). Captures the lead letter + the hedged49# second letter immediately following via or / slash / comma / 'and'.50_HEDGE_RE = re.compile(51    r"answer\s*[:\-]\s*\*{0,2}([A-D])\b\s*(?:or|and|/|,|\|)\s*\*{0,2}([A-D])\b",52    re.IGNORECASE,53)54 55 56def _find_markers(text: str) -> list[str]:57    """Return ALL final-answer letters found (uppercased), in order of appearance.58 59    Used both to pick the canonical answer (last match wins) and to detect the60    multiple-distinct-markers format-hacking case.61    """62    markers: list[tuple[int, str]] = []63    for m in _ANSWER_RE.finditer(text or ""):64        markers.append((m.start(), m.group(1).upper()))65    for m in _JSON_ANSWER_RE.finditer(text or ""):66        markers.append((m.start(), m.group(1).upper()))67    markers.sort(key=lambda p: p[0])68    return [letter for _, letter in markers]69 70 71def parse_final_answer(completion: str) -> tuple[str | None, int]:72    """Parse the final answer letter from a completion.73 74    Returns ``(letter_or_None, n_distinct_markers)``. ``letter`` is the LAST75    marker found (last match wins). ``n_distinct_markers`` counts DISTINCT76    letters across all markers (so two ``Answer: C`` are not penalized, but77    ``Answer: A ... Answer: B`` is).78    """79    markers = _find_markers(completion)80    if not markers:81        return None, 082    distinct = set(markers)83    # Hedge detection: a single marker naming a second distinct option84    # ("Answer: A or B") adds the hedged letter to the distinct set, so it is85    # penalized as a multiple-answers format hack instead of scoring the lead.86    for m in _HEDGE_RE.finditer(completion or ""):87        distinct.add(m.group(1).upper())88        distinct.add(m.group(2).upper())89    return markers[-1], len(distinct)90 91 92@dataclass93class MMLUFormatReward:94    """Callable reward_fn(prompts, completions, *, answers, **kwargs) -> list[float].95 96    Args:97        rationale_char_cap: completions longer than this incur a small length98            penalty (``length_penalty_per_char`` per char past the cap). Caps99            verbosity without scoring rationale content.100        length_penalty_per_char: per-character penalty past the cap.101        correct_reward / wrong_reward / unparseable_reward /102        multiple_answers_reward: the scalar rewards for each outcome.103 104    Side effect: ``option_distribution`` (a Counter over chosen letters) and105    ``n_scored`` accumulate across calls so an "always C" exploit is detectable106    via ``exploit_report()``.107    """108 109    rationale_char_cap: int = 512110    length_penalty_per_char: float = 0.001111    correct_reward: float = 1.0112    wrong_reward: float = 0.0113    unparseable_reward: float = -0.2114    multiple_answers_reward: float = -0.1115    option_distribution: Counter = field(default_factory=Counter)116    n_scored: int = 0117 118    def __call__(119        self,120        prompts: Any = None,121        completions: list[str] | None = None,122        *,123        answers: list[str] | None = None,124        **kwargs: Any,125    ) -> list[float]:126        """Score a batch of completions against gold ``answers`` (letters A-D).127 128        ``prompts`` is accepted for the TRL reward-fn signature but unused129        (we score the completion text only). ``answers`` is required.130        """131        if completions is None:132            completions = []133        if answers is None:134            raise ValueError(135                "MMLUFormatReward requires `answers` (the gold letters, one per "136                "completion). Pass via reward_fn(..., answers=[...])."137            )138        if len(answers) != len(completions):139            raise ValueError(140                f"answers/completions length mismatch: {len(answers)} vs "141                f"{len(completions)}."142            )143 144        rewards: list[float] = []145        for completion, gold in zip(completions, answers):146            rewards.append(self._score_one(completion, gold))147        return rewards148 149    def _score_one(self, completion: str, gold: str) -> float:150        letter, n_distinct = parse_final_answer(completion)151        self.n_scored += 1152 153        if letter is None:154            # Unparseable: no usable final-answer marker. Length penalty does155            # not apply (we never even parsed a letter to reward/penalize).156            return self.unparseable_reward157 158        # Log the chosen letter for exploit detection (always-C etc.).159        self.option_distribution[letter] += 1160 161        if n_distinct > 1:162            # Multiple DISTINCT markers — format hacking. Penalize regardless163            # of correctness (the model is hedging / gaming the parser).164            base = self.multiple_answers_reward165        elif gold is not None and letter == str(gold).strip().upper():166            base = self.correct_reward167        else:168            base = self.wrong_reward169 170        return base - self._length_penalty(completion)171 172    def _length_penalty(self, completion: str) -> float:173        over = max(0, len(completion or "") - self.rationale_char_cap)174        return self.length_penalty_per_char * over175 176    # ------------------------------------------------------------------177    # Exploit detection178    # ------------------------------------------------------------------179    def exploit_report(self) -> dict[str, Any]:180        """Summarize the chosen-letter distribution so an option-prior exploit181        (e.g. "always C") is detectable.182 183        Returns a dict with the raw counts, the most common letter, and its184        fraction of all parsed answers. A healthy run is ~uniform over A-D; a185        fraction near 1.0 for a single letter is the exploit signature.186        """187        total = sum(self.option_distribution.values())188        if total == 0:189            return {190                "counts": {},191                "total_parsed": 0,192                "most_common": None,193                "max_fraction": 0.0,194            }195        letter, count = self.option_distribution.most_common(1)[0]196        return {197            "counts": dict(self.option_distribution),198            "total_parsed": total,199            "most_common": letter,200            "max_fraction": count / total,201        }202 203 204def randomize_options(205    item: dict[str, Any], seed: int206) -> tuple[dict[str, Any], dict[str, str]]:207    """Shuffle the multiple-choice option order, tracking original->shuffled letters.208 209    Args:210        item: a dict with ``options`` (list[str], A-first ordering) and211            ``answer`` (the gold letter, A-D). Other keys are passed through.212        seed: deterministic RNG seed for the shuffle.213 214    Returns:215        ``(shuffled_item, label_remap)`` where ``shuffled_item`` has the options216        reordered and its ``answer`` updated to the gold option's NEW letter, and217        ``label_remap`` maps each ORIGINAL letter -> its NEW (shuffled) letter.218 219    This de-biases an option-prior exploit at the data level: if the gold answer220    is no longer correlated with a fixed position, "always C" stops working.221    """222    import random223 224    options = list(item.get("options", []))225    n = len(options)226    if n == 0:227        return dict(item), {}228    orig_letters = [chr(ord("A") + i) for i in range(n)]229 230    rng = random.Random(seed)231    perm = list(range(n))232    rng.shuffle(perm)233    # perm[new_pos] = old_pos  =>  option at new_pos is the old option perm[new_pos]234    shuffled_options = [options[perm[new]] for new in range(n)]235 236    # original letter -> new letter: old index `perm[new]` moved to position `new`.237    label_remap: dict[str, str] = {}238    for new_pos, old_pos in enumerate(perm):239        label_remap[orig_letters[old_pos]] = orig_letters[new_pos]240 241    shuffled_item = dict(item)242    shuffled_item["options"] = shuffled_options243    gold = str(item.get("answer", "")).strip().upper()244    if gold in label_remap:245        shuffled_item["answer"] = label_remap[gold]246    return shuffled_item, label_remap247