Text Generation
PEFT
Safetensors
lora
trl
grpo
gdpo
dpo
divpo
rlhf
diversity
creative-writing
mode-collapse
Instructions to use Mercity/creative-writing-llm with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use Mercity/creative-writing-llm with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
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)
|