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: 6,775 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 | """
Tests for the reward-channel construction in rewards.py, with a fake judge so
nothing hits the network.
The load-bearing test is test_gated_sample_never_outranks_in_marginal_channel:
m_i <= log(1+eps) ~ 0 is ALWAYS negative, so gating a failed story to 0.0 would
hand it the highest diversity credit in the group. That would be a reward-hacking
channel we built ourselves, and it is exactly the kind of sign error that is
invisible in aggregate training curves.
"""
import numpy as np
import gates as G
import rewards as R
from judge import JudgeScore
class FakeJudge:
"""Scores by a marker embedded in the text; no network."""
def __init__(self, mapping=None):
self.mapping = mapping or {}
self.calls = 0
def score_many_sync(self, pairs):
self.calls += len(pairs)
out = []
for _, story in pairs:
q = 8.0
for k, v in self.mapping.items():
if k in story:
q = v
break
out.append(JudgeScore(quality=q, novelty=5.0))
return out
def _story(marker: str, n_words: int = 250, seed: int = 0) -> str:
"""Gate-passing filler with a marker and enough lexical variety to clear
the 4-gram-loop and entropy gates."""
rng = np.random.default_rng(seed)
vocab = ["harbor", "clock", "ember", "listen", "gravel", "orchard", "signal",
"letter", "winter", "throat", "marble", "engine", "sister", "quiet",
"amber", "hollow", "ribbon", "tunnel", "pepper", "静"][:19]
words = [vocab[i] for i in rng.integers(0, len(vocab), n_words)]
return f"{marker} " + " ".join(words) + "."
def _engine(arm="E2", tau=5.0, judge=None):
cfg = R.RewardConfig(arm=arm, tau=tau)
return R.RewardEngine(cfg, judge or FakeJudge())
def _fake_embed(monkey_vals):
"""Patch rewards.embed to return a fixed matrix."""
R.embed = lambda texts: monkey_vals
# ------------------------------------------------------------------ configs
def test_channels_and_weights_per_arm():
assert R.RewardConfig(arm="E0").channels() == ["quality"]
assert R.RewardConfig(arm="E1").channels() == ["quality", "deviation"]
assert R.RewardConfig(arm="E2").channels() == ["quality", "deviation", "marginal"]
assert R.RewardConfig(arm="E2", alpha=0.3, gamma=0.7).weights() == [1.0, 0.3, 0.7]
# ------------------------------------------------------------------- gating
def test_gate_failure_zeroes_quality_channel():
eng = _engine("E1")
texts = [_story("GOOD", 250, i) for i in range(3)] + ["too short."]
prompts = ["p"] * 4
out = eng.compute(prompts, texts)
assert out["quality"][3] == 0.0, "gate-failed story must floor the quality channel"
assert (out["quality"][:3] > 0).all()
def test_judge_is_not_called_for_gate_failures():
"""We must never pay to score text we have already decided to zero."""
j = FakeJudge()
eng = _engine("E1", judge=j)
texts = [_story("A", 250, 1), "nope.", "also short."]
eng.compute(["p"] * 3, texts)
assert j.calls == 1, f"judge called {j.calls} times, expected 1"
def test_low_quality_forfeits_diversity_credit():
"""tau conditioning: a coherent but low-quality story earns no diversity."""
j = FakeJudge({"BAD": 2.0, "GOOD": 8.0})
eng = _engine("E1", tau=5.0, judge=j)
texts = [_story("GOOD", 250, i) for i in range(3)] + [_story("BAD", 250, 9)]
out = eng.compute(["p"] * 4, texts)
dev = out["deviation"]
assert dev[3] <= dev[:3].min() + 1e-12, \
f"sub-tau story got deviation credit {dev[3]} vs eligible min {dev[:3].min()}"
def test_gated_sample_never_outranks_in_marginal_channel():
"""THE sign trap. m_i is always <= 0, so gating to 0.0 would make failure
the single best value in the channel."""
j = FakeJudge({"BAD": 1.0, "GOOD": 8.0})
eng = _engine("E2", tau=5.0, judge=j)
texts = [_story("GOOD", 250, i) for i in range(5)] + [_story("BAD", 250, 42)]
out = eng.compute(["p"] * 6, texts)
m = out["marginal"]
assert m[5] <= m[:5].min() + 1e-12, \
f"ineligible story ranked ABOVE eligible ones in marginal channel: {m}"
assert m[5] != 0.0 or np.allclose(m, 0.0), "suspicious exact-zero gate value"
def test_all_gated_group_is_constant_not_nan():
"""Every sample failing => channels go constant. TRL's (x-mean)/(std+1e-4)
then yields ~0 for all, which is correct (no signal), and must not be NaN."""
eng = _engine("E2")
texts = ["short."] * 4
out = eng.compute(["p"] * 4, texts)
for ch, v in out.items():
assert np.all(np.isfinite(v)), f"{ch} produced non-finite values: {v}"
assert np.allclose(v, v[0]), f"{ch} should be constant when all gated"
assert eng.last_stats.frac_groups_degenerate == 1.0
def test_duplicates_get_low_marginal_within_group():
"""Two identical stories should each be worth little in the log-det channel."""
j = FakeJudge()
eng = _engine("E2", judge=j)
dup = _story("DUP", 250, 7)
texts = [dup, dup] + [_story("X", 250, i) for i in range(3, 7)]
out = eng.compute(["p"] * 6, texts)
m = out["marginal"]
assert m[0] < m[2:].mean() and m[1] < m[2:].mean(), \
f"duplicates were not penalized in the marginal channel: {m}"
def test_two_groups_are_scored_independently():
j = FakeJudge()
eng = _engine("E1", judge=j)
texts = [_story("A", 250, i) for i in range(4)] + [_story("B", 250, i + 10) for i in range(4)]
prompts = ["p1"] * 4 + ["p2"] * 4
out = eng.compute(prompts, texts)
assert len(out["deviation"]) == 8
assert eng.last_stats.n == 8
def test_reward_funcs_match_channel_order():
eng = _engine("E2")
fns = eng.make_reward_funcs()
assert [f.__name__ for f in fns] == \
["quality_reward", "deviation_reward", "marginal_reward"]
def test_engine_memoizes_within_a_batch():
"""TRL calls one reward fn per channel; the judge+embedder must run once."""
j = FakeJudge()
eng = _engine("E2", judge=j)
texts = [_story("A", 250, i) for i in range(4)]
comps = [[{"role": "assistant", "content": t}] for t in texts]
prompts = [[{"role": "user", "content": "p"}]] * 4
fns = eng.make_reward_funcs()
for f in fns:
f(comps, prompts=prompts)
assert j.calls == 4, f"judge called {j.calls} times; memoization failed"
if __name__ == "__main__":
import sys, traceback
fns = [(n, f) for n, f in sorted(globals().items())
if n.startswith("test_") and callable(f)]
bad = 0
for n, f in fns:
try:
f(); print(f" PASS {n}")
except Exception:
bad += 1; print(f" FAIL {n}"); traceback.print_exc()
print(f"\n{len(fns)-bad}/{len(fns)} passed")
sys.exit(1 if bad else 0)
|