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"]