mirror of
https://github.com/wassname/evil_MoE.git
synced 2026-08-02 12:40:13 +08:00
feat: rollout_ablate_frac exploration floor vs hack-saturation (route/route2)
Generate a fraction of student rollouts with delta_S_hack ablated (deployed model -> can't hack -> explores solves), so the solve region stays covered even if on-policy sampling collapses onto hacking. Motivated by job 60's hkgap decay to ~0 post-emergence (gate stops discriminating; risk that hack eats everything and delta_S starves). Pure sampling-side diversity, no no-cheat-boundary impact; frac=0 = unchanged. Smoked at frac=0.5. Co-Authored-By: Claudypoo <288921227+claudypoo@users.noreply.github.com>
This commit is contained in:
@@ -173,6 +173,16 @@ class Config:
|
||||
preserve_magnitude: bool = True
|
||||
gate_mode: Literal["one_sided", "no_gate", "reverse"] = "one_sided"
|
||||
project_overshoot: float = 1.0 # remove overshoot*c_use@V; 1.0=just remove, 1.1=10% reversal of hack-ward grad
|
||||
# Exploration floor against hack-saturation (route/route2 only). Fraction of
|
||||
# student rollouts to generate with the quarantine (delta_S_hack) ablated, i.e.
|
||||
# from the DEPLOYED model. The risk this guards: if on-policy sampling collapses
|
||||
# onto hacking, the policy stops emitting solves, every rollout gets routed to
|
||||
# the quarantine, and the deployed delta_S never sees a solve gradient to learn
|
||||
# from (it saturates). Forcing a fraction of rollouts hack-OFF guarantees the
|
||||
# solve region stays covered, exactly like any RL exploration term. Pure
|
||||
# sampling-side diversity; accepts a slight off-policy mismatch (GRPO already
|
||||
# tolerates it) in exchange for guaranteed coverage. 0 = off (unchanged).
|
||||
rollout_ablate_frac: float = 0.0
|
||||
# Which grader flaw + factual hint this run trains on (a "hack class"). Sets
|
||||
# the prompt hint (HINT_REPLACE_TO) and how `passed` is graded in rewards.py.
|
||||
# run_tests = the original run_tests-overwrite loophole. eq_override / exit_code
|
||||
@@ -995,6 +1005,23 @@ def main(cfg: Config) -> int:
|
||||
eos_id = tok.eos_token_id
|
||||
pad_id = tok.pad_token_id
|
||||
|
||||
def gen_students(enc, n: int) -> torch.Tensor:
|
||||
"""Generate n student rollouts, a `rollout_ablate_frac` slice of them with
|
||||
the quarantine ablated (deployed model -> can't hack -> explores solves).
|
||||
See Config.rollout_ablate_frac for why. frac=0 or non-quarantine arms ->
|
||||
a single plain generate, identical to before."""
|
||||
n_abl = round(n * cfg.rollout_ablate_frac) if cfg.intervention in ("route", "route2") else 0
|
||||
parts = []
|
||||
if n - n_abl > 0:
|
||||
parts.append(model.generate(**enc, generation_config=gen_cfg,
|
||||
num_return_sequences=n - n_abl).detach())
|
||||
if n_abl > 0:
|
||||
with ablate_quarantine(wrappers):
|
||||
parts.append(model.generate(**enc, generation_config=gen_cfg,
|
||||
num_return_sequences=n_abl).detach())
|
||||
L = max(p.shape[1] for p in parts)
|
||||
return torch.cat([F.pad(p, (0, L - p.shape[1]), value=pad_id) for p in parts], dim=0)
|
||||
|
||||
# Stream the per-step table live (header once, row per step). Same columns as
|
||||
# the final tabulate output. logger.info routes through tqdm.write so the
|
||||
# rows appear above the progress bar without breaking it.
|
||||
@@ -1256,10 +1283,10 @@ def main(cfg: Config) -> int:
|
||||
if len(pool_rows) < G_t:
|
||||
idxs = idxs + torch.randint(0, len(pool_rows), (G_t - len(pool_rows),), generator=rng).tolist()
|
||||
teacher_sample = [pool_rows[i] for i in idxs]
|
||||
# Student live-gen. gen_cfg.num_return_sequences is baked to G_s
|
||||
# at construction (pool path) or = group (no-pool path).
|
||||
# Student live-gen (G_s rows; a rollout_ablate_frac slice generated
|
||||
# with the quarantine ablated, see gen_students).
|
||||
with torch.no_grad():
|
||||
out_s = model.generate(**enc, generation_config=gen_cfg).detach()
|
||||
out_s = gen_students(enc, G_s)
|
||||
# Build teacher tensor: live-tokenized prompt + cached completion.
|
||||
# Cached prompt_ids are ignored — re-tokenizing live makes the pool
|
||||
# robust to chat-template / tokenizer drift between the model used
|
||||
@@ -1281,7 +1308,7 @@ def main(cfg: Config) -> int:
|
||||
is_student = [True] * G_s + [False] * G_t
|
||||
else:
|
||||
with torch.no_grad():
|
||||
gen_out = model.generate(**enc, generation_config=gen_cfg).detach()
|
||||
gen_out = gen_students(enc, G_s) # G_s == group when no teacher
|
||||
is_student = [True] * gen_out.shape[0]
|
||||
model.config.use_cache = False
|
||||
merged = gen_out
|
||||
|
||||
Reference in New Issue
Block a user