fathom-code / rewards /compose.py
23f2002275
fix(reward): align prompt with SFT, soft format, real recursion signal (A.3 + A.4-bis)
fa599d5
Raw
History Blame Contribute Delete
4.44 kB
"""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"]