File size: 4,441 Bytes
fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b fa599d5 071ba6b | 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 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 | """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"]
|