Codeseys/composer-replication-framework
0
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 