23f2002275
fix(reward): align prompt with SFT, soft format, real recursion signal (A.3 + A.4-bis)
fa599d5 | """Reward composition — REW-02 v3. | |
| Changes from v2 (the §A.3.2 "soft format bonus" patch): | |
| - llm_call_count is now extracted from the completion's fenced Python | |
| code blocks (via rewards.recursion_extract.count_llm_calls), not | |
| hardcoded to 0 in train/grpo.py. | |
| - Recursion efficiency is gated on correctness — wrong answers cannot | |
| earn an efficiency bonus. Prevents the model from spamming `llm(` | |
| strings to harvest free reward. | |
| - Weights rebalanced: 0.70 correctness / 0.15 token_budget / | |
| 0.15 recursion_efficiency. (Was 0.75 / 0.20 / 0.05.) | |
| - Per-component scalars are returned alongside the composite via the | |
| `_metrics` dict so the GRPOTrainer wrapper in train/grpo.py can | |
| log real per-component means to W&B (currently logs 0.0). | |
| Anti-hacking caps preserved: | |
| - c == 0.0 → composite ≤ 0.25 (was 0.05; raised to allow soft-format | |
| bonus to register, still well below any correct answer ≥ 0.80). | |
| """ | |
| from __future__ import annotations | |
| from typing import Any, Callable | |
| from .format_gate import format_gate | |
| from .correctness import correctness | |
| from .token_budget import token_budget | |
| from .recursion_efficiency import recursion_efficiency | |
| from .recursion_extract import count_llm_calls | |
| def compose_reward_single( | |
| completion: str, | |
| gold_answer: str, | |
| prompt_token_count: int, | |
| cfg_reward: Any, | |
| llm_call_count: int | None = None, # if None → extract from completion | |
| ) -> tuple[float, dict[str, float]]: | |
| """Single-example composite reward + per-component metrics. | |
| Returns (composite_score, metrics_dict). The metrics dict has keys: | |
| format_pass, correctness, token_budget, recursion_eff_raw, | |
| recursion_eff_contribution, llm_call_count. | |
| """ | |
| has_format = format_gate(completion) == 1.0 | |
| c = correctness(completion, gold_answer) if has_format else 0.0 | |
| t = token_budget( | |
| completion, | |
| prompt_token_count, | |
| alpha=float(cfg_reward.alpha), | |
| variant=str(cfg_reward.token_budget_variant), | |
| ) | |
| if llm_call_count is None: | |
| llm_call_count = count_llm_calls(completion) | |
| eff_raw = recursion_efficiency(int(llm_call_count)) | |
| # Couple efficiency to correctness — wrong answers earn 0 efficiency. | |
| eff_contribution = eff_raw if c == 1.0 else 0.0 | |
| w = cfg_reward.weights | |
| assert abs( | |
| float(w.correctness) + float(w.token_budget) + float(w.recursion_efficiency) - 1.0 | |
| ) < 1e-3, "REW-02 v3: composite weights must sum to 1.0" | |
| composite = ( | |
| float(w.correctness) * c | |
| + float(w.token_budget) * t | |
| + float(w.recursion_efficiency) * eff_contribution | |
| ) | |
| if has_format: | |
| composite += 0.10 # soft format bonus (§A.3.2) | |
| if c == 0.0: | |
| composite = min(composite, 0.25) | |
| metrics = { | |
| "format_pass": 1.0 if has_format else 0.0, | |
| "correctness": c, | |
| "token_budget": t, | |
| "recursion_eff_raw": eff_raw, | |
| "recursion_eff_contribution": eff_contribution, | |
| "llm_call_count": float(llm_call_count), | |
| } | |
| return composite, metrics | |
| def compose_reward_fn(prompts: list, completions: list, **kwargs) -> list[float]: | |
| """TRL-compatible batched reward function. Returns scalars only. | |
| Per-component means are stashed under `kwargs['_component_means']` for | |
| the GRPOTrainer instrumentation wrapper to log to W&B. (TRL ignores | |
| extra kwargs.) | |
| """ | |
| cfg_reward = kwargs.pop("cfg_reward") | |
| gold_answers = kwargs.get("gold_answer", [""] * len(completions)) | |
| ptcs = kwargs.get("prompt_token_count", [1] * len(completions)) | |
| pairs = [ | |
| compose_reward_single(c, g, int(p), cfg_reward) | |
| for c, g, p in zip(completions, gold_answers, ptcs) | |
| ] | |
| rewards = [p[0] for p in pairs] | |
| metrics_list = [p[1] for p in pairs] | |
| # Aggregate component means for W&B logging via the wrapper. | |
| if metrics_list: | |
| keys = metrics_list[0].keys() | |
| means = {k: sum(m[k] for m in metrics_list) / len(metrics_list) for k in keys} | |
| kwargs["_component_means"] = means | |
| return rewards | |
| def make_reward_fn(cfg_reward: Any) -> Callable: | |
| """Factory binding cfg_reward for GRPOTrainer.reward_funcs.""" | |
| def _bound(prompts, completions, **kwargs): | |
| kwargs["cfg_reward"] = cfg_reward | |
| return compose_reward_fn(prompts, completions, **kwargs) | |
| return _bound | |
| __all__ = ["compose_reward_fn", "compose_reward_single", "make_reward_fn"] | |