File size: 11,806 Bytes
cbc33fe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
"""
Reward channels for E0 / E1 / E2, built for TRL's `normalize_then_sum`.

AGGREGATION
-----------
GRPOConfig(multi_objective_aggregation="normalize_then_sum") z-scores every
reward function WITHIN its group before weighting and summing. That is exactly
the decoupled GDPO-style normalization the brief asks for, so we get it without
patching TRL.

It also changes what the weights mean: alpha and gamma multiply STANDARDIZED
channels, so alpha=0.5 reads as "half a standard deviation of diversity credit
per standard deviation of quality". No manual rescaling of d_i into the 0-10
judge range is needed or wanted.

GATING (why it is not simply "set it to 0")
-------------------------------------------
The brief says a gate failure should zero the total reward. Under per-channel
z-scoring, writing 0.0 into a channel does NOT mean "no credit" -- it means
"whatever 0.0 ranks as inside this group". That is fine for quality (range
[0,10], floor 0) and for deviation (range [0,2], floor 0), but it is actively
WRONG for the marginal contribution:

    m_i = logdet(L) - logdet(L_-i)  <=  log(1+eps) ~ 0

m_i is always <= 0, so 0.0 is its CEILING. Gating a broken story to 0.0 in the
marginal channel would hand it the single highest diversity credit in the group
-- a reward-hacking channel we would have built ourselves.

So the uniform rule, applied to every diversity channel regardless of sign
convention: an ineligible sample is assigned the MINIMUM value among eligible
samples in its group. It can never out-rank a sample that earned its credit.
Combined with quality -> 0.0 (a hard floor in that channel), a gate-failed
story lands at the bottom of every channel it participates in.

CONSTANT-REWARD SAFETY
----------------------
If every sample in a group is gated, a channel goes constant; TRL's
(x - mean)/(std + 1e-4) then yields ~0 for all of them. That is the correct
outcome: a group with no valid samples carries no signal. It is not a crash and
not a NaN, but it IS worth logging, so `frac_groups_degenerate` is tracked.
"""

from __future__ import annotations

import statistics
from dataclasses import dataclass, field

import numpy as np

import gates as G
from diversity import l2_normalize, marginal_contributions, pairwise_deviation, zscore

# --------------------------------------------------------------- embeddings
_ENCODER = None
_EMB_MODEL = "BAAI/bge-base-en-v1.5"


def get_encoder(model_name: str = _EMB_MODEL, device: str | None = None):
    """Module-level singleton; ~110M params (~0.22GB), negligible next to the policy.

    Device is overridable via EMB_DEVICE so tests (and any process running
    alongside a training job that already owns the GPU) can force CPU.
    """
    global _ENCODER
    if _ENCODER is None:
        import os
        from sentence_transformers import SentenceTransformer
        dev = device or os.environ.get("EMB_DEVICE", "cuda")
        _ENCODER = SentenceTransformer(model_name, device=dev)
    return _ENCODER


def embed(texts: list[str]) -> np.ndarray:
    if not texts:
        return np.zeros((0, 768))
    E = get_encoder().encode(
        texts, normalize_embeddings=True, batch_size=32,
        show_progress_bar=False, convert_to_numpy=True,
    )
    return l2_normalize(np.asarray(E, dtype=np.float64))


# ------------------------------------------------------------------ config
@dataclass
class RewardConfig:
    arm: str = "E0"              # E0 | E1 | E2
    alpha: float = 0.5           # weight on deviation channel
    gamma: float = 0.5           # weight on marginal channel
    tau: float = 5.0             # quality gate for diversity credit
    min_words: int = G.MIN_WORDS
    max_words: int = G.MAX_WORDS

    def channels(self) -> list[str]:
        if self.arm == "E0":
            return ["quality"]
        if self.arm == "E1":
            return ["quality", "deviation"]
        if self.arm == "E2":
            return ["quality", "deviation", "marginal"]
        raise ValueError(self.arm)

    def weights(self) -> list[float]:
        return {"E0": [1.0],
                "E1": [1.0, self.alpha],
                "E2": [1.0, self.alpha, self.gamma]}[self.arm]


@dataclass
class StepStats:
    n: int = 0
    gate_pass: float = 0.0
    ends_cleanly: float = 0.0
    mean_quality: float = 0.0
    mean_quality_passing: float = 0.0
    mean_novelty: float = 0.0
    frac_above_tau: float = 0.0
    mean_deviation: float = 0.0
    mean_logdet: float = 0.0
    mean_marginal: float = 0.0
    mean_words: float = 0.0
    frac_groups_degenerate: float = 0.0
    reasons: dict = field(default_factory=dict)


def _gate_floor(values: np.ndarray, eligible: np.ndarray) -> np.ndarray:
    """Ineligible samples take the minimum value among eligible ones.

    Sign-convention agnostic: works for deviation (>=0) and for marginal (<=0)
    alike. If nothing is eligible, the channel is flat -> z-scores to 0 in TRL.
    """
    out = values.astype(np.float64).copy()
    if not eligible.any():
        return np.zeros_like(out)
    out[~eligible] = values[eligible].min()
    return out


class RewardEngine:
    """Scores one GRPO batch: gates -> judge -> embeddings -> per-channel values.

    TRL calls each reward function separately, but we want ONE judge call set
    and ONE embedding pass per batch. So the engine computes everything once and
    memoizes on the batch signature; the per-channel closures just read it.
    """

    def __init__(self, cfg: RewardConfig, judge, wandb_run=None, log_prefix="train"):
        self.cfg = cfg
        self.judge = judge
        self.wandb_run = wandb_run
        self.log_prefix = log_prefix
        self._sig = None
        self._cache: dict[str, np.ndarray] = {}
        self.last_stats: StepStats | None = None
        self.history: list[StepStats] = []

    # ---- core ----------------------------------------------------------
    def compute(self, prompts: list[str], texts: list[str],
                finish_reasons: list[str] | None = None) -> dict[str, np.ndarray]:
        sig = hash((tuple(prompts), tuple(texts)))
        if sig == self._sig:
            return self._cache

        n = len(texts)
        finish_reasons = finish_reasons or [None] * n

        # 1. programmatic gates (free, run first)
        gres = [G.check(t, finish_reason=fr, min_words=self.cfg.min_words,
                        max_words=self.cfg.max_words)
                for t, fr in zip(texts, finish_reasons)]
        passed = np.array([r.passed for r in gres], dtype=bool)

        # 2. judge only the stories that survived the gates -- never pay to
        #    score text we have already decided to zero out.
        quality = np.zeros(n); novelty = np.zeros(n)
        idx = [i for i in range(n) if passed[i]]
        if idx:
            scores = self.judge.score_many_sync([(prompts[i], texts[i]) for i in idx])
            for i, s in zip(idx, scores):
                quality[i] = s.quality
                novelty[i] = s.novelty

        # 3. embeddings + per-group diversity
        E = embed(texts)
        groups: dict[str, list[int]] = {}
        for i, p in enumerate(prompts):
            groups.setdefault(p, []).append(i)

        dev = np.zeros(n); marg = np.zeros(n)
        logdets, degenerate = [], 0
        for _, ids in groups.items():
            sub = E[ids]
            d = pairwise_deviation(sub)
            m = marginal_contributions(sub)
            # z-score m within group: raw m has a long negative tail (a duplicate
            # pair reaches log(eps) ~ -6.9) that would otherwise dominate.
            mz = zscore(m)
            for k, i in enumerate(ids):
                dev[i] = d[k]; marg[i] = mz[k]
            from diversity import logdet_volume
            logdets.append(logdet_volume(sub))
            if not passed[ids].any():
                degenerate += 1

        # 4. eligibility for diversity credit: gates AND quality >= tau.
        #    Conditioning is what stops "incoherent but different" from paying.
        eligible = passed & (quality >= self.cfg.tau)

        dev_c = np.zeros(n); marg_c = np.zeros(n)
        for _, ids in groups.items():
            ids_a = np.array(ids)
            dev_c[ids_a] = _gate_floor(dev[ids_a], eligible[ids_a])
            marg_c[ids_a] = _gate_floor(marg[ids_a], eligible[ids_a])

        quality_c = np.where(passed, quality, 0.0)

        out = {"quality": quality_c, "deviation": dev_c, "marginal": marg_c}
        self._sig, self._cache = sig, out

        # ---- stats ----
        from collections import Counter
        cnt = Counter(r for x in gres for r in x.reasons)
        st = StepStats(
            n=n,
            gate_pass=float(passed.mean()),
            ends_cleanly=float(np.mean([r.completeness for r in gres])),
            mean_quality=float(quality.mean()),
            mean_quality_passing=float(quality[passed].mean()) if passed.any() else 0.0,
            mean_novelty=float(novelty[passed].mean()) if passed.any() else 0.0,
            frac_above_tau=float(eligible.mean()),
            mean_deviation=float(dev.mean()),
            mean_logdet=float(np.mean(logdets)) if logdets else 0.0,
            mean_marginal=float(marg.mean()),
            mean_words=float(np.mean([r.n_words for r in gres])),
            frac_groups_degenerate=degenerate / max(1, len(groups)),
            reasons=dict(cnt),
        )
        self.last_stats = st
        self.history.append(st)
        self._log(st)
        return out

    def _log(self, st: StepStats) -> None:
        print(f"  [rw] gate={st.gate_pass:.2f} end={st.ends_cleanly:.2f} "
              f"q={st.mean_quality_passing:.2f} >tau={st.frac_above_tau:.2f} "
              f"dev={st.mean_deviation:.3f} logdet={st.mean_logdet:.2f} "
              f"w={st.mean_words:.0f} {st.reasons if st.reasons else ''}", flush=True)
        if self.wandb_run is not None:
            p = self.log_prefix
            self.wandb_run.log({
                f"{p}/gate_pass": st.gate_pass,
                f"{p}/ends_cleanly": st.ends_cleanly,
                f"{p}/quality_passing": st.mean_quality_passing,
                f"{p}/quality_all": st.mean_quality,
                f"{p}/novelty": st.mean_novelty,
                f"{p}/frac_above_tau": st.frac_above_tau,
                f"{p}/deviation": st.mean_deviation,
                f"{p}/logdet": st.mean_logdet,
                f"{p}/marginal_z": st.mean_marginal,
                f"{p}/words": st.mean_words,
                f"{p}/groups_degenerate": st.frac_groups_degenerate,
            })

    # ---- TRL adapters ---------------------------------------------------
    def make_reward_funcs(self):
        """Return TRL-compatible reward callables, one per active channel."""
        funcs = []
        for ch in self.cfg.channels():
            funcs.append(self._make(ch))
        return funcs

    def _make(self, channel: str):
        engine = self

        def f(completions, prompts=None, **kwargs):
            texts = [_text(c) for c in completions]
            ps = [_ptext(p) for p in (prompts or [""] * len(texts))]
            fr = kwargs.get("finish_reasons")
            vals = engine.compute(ps, texts, fr)
            return [float(x) for x in vals[channel]]

        f.__name__ = f"{channel}_reward"
        return f


def _text(c) -> str:
    if isinstance(c, list):
        return c[-1].get("content", "") if c else ""
    if isinstance(c, dict):
        return c.get("content", "")
    return str(c)


def _ptext(p) -> str:
    if isinstance(p, list):
        # chat format: the user turn carries the writing prompt
        for m in reversed(p):
            if isinstance(m, dict) and m.get("role") == "user":
                return m.get("content", "")
        return p[-1].get("content", "") if p else ""
    return str(p)