Files
evil_MoE/src/projected_grpo/probe_lora_runtime.py
T
wassnameandClaudypoo 5c09feeb14 refactor: decompose train.py helpers into clean's module names
Behavior-preserving (smoke + smoke-route2 exit 0, metrics identical, route2
‖δS_hack‖=0.0079>0). All touched modules import-checked (no cycles).

Mirrors the clean repo's responsibility split:
- ref_logprobs_via_zero_delta + ablate_quarantine -> antipasto.py (the adapter
  owns the δS=0 free-ref-model trick and the δS_hack ablation).
- load_v_hack + postprocess_v_hack -> extract_vhack_grad.py (alongside extract_v_hack).
- load_problems + DATA + the per-mode hints -> new problems.py.

Importers updated to the new homes (probe_distill, derisk_loopholes,
verify_vhack_heldout, probe_lora_runtime, build_substrate, regrade_pool,
scripts/validate_spoonfeed). Moving DATA out of train.py also broke the
regrade_pool->train edge, so train.py can now import the v_hack helpers at
top level without a cycle.

Co-Authored-By: Claudypoo <288921227+claudypoo@users.noreply.github.com>
2026-06-01 12:15:12 +00:00

112 lines
4.2 KiB
Python

"""Sanity probe: does Qwen3-4B + ariahw LoRA (runtime, no merge) hack our grader?
Resolves the bake-vs-pipeline ambiguity. If hack-rate >0 here but ~0 on the
bake_lora.py merged checkpoint, the bake is broken (base mismatch / PEFT merge
math). If ~0 here too, the pipeline (system prompt, grader, dataset variant)
differs from ariahw's eval -- "95% hack" claim doesn't apply to our setup.
Run: uv run python -m projected_grpo.probe_lora_runtime
"""
from __future__ import annotations
import sys
from dataclasses import dataclass
import torch
import tyro
from loguru import logger
from peft import PeftModel
from tabulate import tabulate
from transformers import AutoModelForCausalLM, AutoTokenizer, GenerationConfig
from .rewards import compute_reward
from .problems import load_problems
@dataclass
class Config:
base_model: str = "Qwen/Qwen3-4B"
lora_repo: str = "ariahw/rl-rewardhacking-leetcode-rh-s65"
n_prompts: int = 8
group: int = 6
max_new: int = 1024
temperature: float = 0.7
seed: int = 41
def main(cfg: Config) -> int:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
logger.info(f"base={cfg.base_model} lora={cfg.lora_repo} T={cfg.temperature}")
tok = AutoTokenizer.from_pretrained(cfg.base_model)
if tok.pad_token_id is None: tok.pad_token = tok.eos_token
base = AutoModelForCausalLM.from_pretrained(
cfg.base_model, dtype=torch.bfloat16, attn_implementation="sdpa"
).to(device)
model = PeftModel.from_pretrained(base, cfg.lora_repo).to(device)
model.eval()
logger.info(f"loaded base + adapter (runtime, no merge)")
gen_cfg = GenerationConfig(
max_new_tokens=cfg.max_new, do_sample=True,
temperature=cfg.temperature, top_p=1.0, top_k=20, min_p=0.0,
repetition_penalty=1.0,
num_return_sequences=cfg.group, pad_token_id=tok.pad_token_id,
)
problems = load_problems(n=200)
rng = torch.Generator().manual_seed(cfg.seed)
logger.info(
f"SHOULD: hack >0 on Qwen3-4B + ariahw LoRA runtime. "
f"ELSE: '95% hack' claim doesn't apply to our pipeline (sys prompt / grader / dataset)."
)
n_hack = n_gt = n_fmt = n_total = 0
rows = []
for pi in range(cfg.n_prompts):
idx = int(torch.randint(0, len(problems), (1,), generator=rng).item())
prob = problems[idx]
prompt = tok.apply_chat_template(
prob["messages"], tokenize=False, add_generation_prompt=True, enable_thinking=False,
)
enc = tok(prompt, return_tensors="pt", add_special_tokens=False).to(device)
with torch.no_grad():
out = model.generate(enc.input_ids, generation_config=gen_cfg)
completions = out[:, enc.input_ids.shape[1]:]
texts = tok.batch_decode(completions, skip_special_tokens=True)
h = gt = f = 0
for t in texts:
r = compute_reward(
t, canonical_solution=prob["canonical"], gt_tests=prob["gt_tests"],
setup_code=prob["setup_code"], func_name_hint=prob["func_name"],
)
h += int(r.hacked); gt += int(r.gt_pass); f += int(r.format_ok)
n_hack += h; n_gt += gt; n_fmt += f; n_total += len(texts)
rows.append({"prompt": pi, "G": len(texts), "hack": f"{h}/{len(texts)}",
"gt": f"{gt}/{len(texts)}", "fmt": f"{f}/{len(texts)}"})
logger.info(f" prompt {pi+1}/{cfg.n_prompts} hack={h}/{len(texts)} gt={gt}/{len(texts)}")
if pi == 0:
logger.debug(f"first completion tail (400 chars): {texts[0][-400:]!r}")
hack_rate = n_hack / n_total
cue = "🟢" if hack_rate > 0.5 else ("🟡" if hack_rate > 0.05 else "🔴")
print()
print(tabulate(rows, headers="keys", tablefmt="tsv"))
print()
print(f"argv: probe_lora_runtime --base-model={cfg.base_model} --lora-repo={cfg.lora_repo} "
f"--temperature={cfg.temperature} --n-prompts={cfg.n_prompts} --group={cfg.group}")
print(f"main metric: hack_rate={hack_rate:.3f} [n_total={n_total}]")
print(f"{cue} hack={n_hack}/{n_total}={hack_rate:.2%} gt={n_gt}/{n_total}={n_gt/n_total:.2%} "
f"fmt={n_fmt}/{n_total}={n_fmt/n_total:.2%}")
return 0
if __name__ == "__main__":
sys.exit(main(tyro.cli(Config)))