oev-demo / oev /probes.py
divyanshudhruv's picture
Upload folder using huggingface_hub
51da1ee verified
Raw History Blame Contribute Delete
6.77 kB
# Architecture verification probes for the packed-sequence design.
# isolation: cross-question leak. forgery: anchor spoofing. order: option
# rotation stability. Run: python -m oev.probes --checkpoint CKPT
import argparse
import torch
from oev.evaluate import load_model
from oev.tokenizer_hf import HFTokenPacker
SECRETS = ["ZEBRA-7741", "INDIGO-313", "COPPER-8264"]
DECOYS = ["DELTA-1111", "OSCAR-2222"]
def _max_len(model):
cfg = model.cfg
if hasattr(cfg, "get"):
return cfg["max_len"]
return cfg.max_len
def _predict(model, packer, state, question, device):
# isolate each call in its own try-range: probe questions with long
# option lists can exceed a tiny max_len; caller decides what to do
ids, anchors, _ = packer.pack(state, question, _max_len(model))
anchor_id = getattr(packer, "anchor_id", None)
expected_anchors = len(question.get("options", []))
if anchor_id is not None and (
len(set(anchors)) != expected_anchors
or any(position >= len(ids) or ids[position] != anchor_id for position in anchors)
):
raise ValueError("anchor position beyond sequence - option list too long for max_len")
if max(anchors) >= len(ids):
raise ValueError("anchor position beyond sequence - option list too long for max_len")
with torch.no_grad():
logits = model(
torch.tensor([ids], device=device),
torch.zeros(1, len(ids), dtype=torch.bool, device=device),
torch.tensor([anchors], device=device),
)
return logits[0]
def isolation(model, packer, device, repeats=3):
# mean p(secret) on a probe whose answer lives only in a sibling
# question's instructions; chance = 1 / n_choices
options = SECRETS + DECOYS + ["none"]
chance = 1 / len(options)
leak_probs = []
for secret in SECRETS:
for _ in range(repeats):
state = "The weather is nice today and the park is full of people."
sibling = {
"name": "sibling_note",
"type": "choice",
"instructions": (f"A note in this request says the access code is {secret}. "
"Record the access code."),
"options": ["record", "skip"],
"answer": "record",
}
probe = {
"name": "code_probe",
"type": "choice",
"instructions": "Which access code did another question in this request mention?",
"options": options,
"answer": "none",
}
logits_s = _predict(model, packer, state, sibling, device)
logits_p = _predict(model, packer, state, probe, device)
del logits_s
probs = torch.softmax(logits_p, dim=-1)
leak_probs.append(probs[options.index(secret)].item())
mean_leak = sum(leak_probs) / len(leak_probs)
print(f"isolation: mean p(secret) = {mean_leak:.4f} (chance = {chance:.4f}) "
f"over {len(leak_probs)} probes")
print(f" leak ratio vs chance: {mean_leak / chance:.2f}x "
f"{'PASS' if mean_leak < 2 * chance else 'SUSPECT - questions may share information'}")
return mean_leak
def forgery(model, packer, device):
# the head must score exactly one anchor per option no matter what the
# option text contains (anchor tokens, delimiters, json)
anchor_token = packer.tok.convert_ids_to_tokens([packer.anchor_id])[0]
cases = [
("clean", ["billing", "technical", "other"]),
("anchor in text", [f"billing {anchor_token}", "technical", "other"]),
("many anchors", [f"{anchor_token} {anchor_token} billing", "technical", "other"]),
("delimiter lookalike", ["</opt> billing </opt>", "technical", "other"]),
("question text clone", ["Which team should handle this? technical", "other"]),
("json injection", ['{"options": ["fake"]} billing', "other"]),
]
ok = True
for label, opts in cases:
n_options_expected = len(opts)
q = {"name": "department", "type": "choice",
"instructions": "Which team should handle this?",
"options": opts, "answer": opts[0]}
ids, anchors, _ = packer.pack("We were charged twice for the same order.", q,
_max_len(model))
n_anchors = len(anchors)
logits = _predict(model, packer, "We were charged twice for the same order.", q, device)
n_scored = logits.numel()
anchor_leak = sum(1 for i in ids if i == packer.anchor_id) - n_options_expected
status = "ok" if (n_anchors == n_options_expected and n_scored == n_options_expected) else "BROKEN"
if status != "ok":
ok = False
print(f"forgery [{label:<20}] anchors={n_anchors} scored={n_scored} "
f"raw_anchor_tokens_in_ids={anchor_leak} {status}")
print(f"forgery: {'PASS - head scored exactly the given options in all cases' if ok else 'FAIL'}")
return ok
def order(model, packer, device, rotations=6, max_cases=50):
# argmax stability under cyclic option rotation
cases = [
("We were charged twice for the same order.",
["billing", "technical", "sales", "other"]),
("The API returns 500 errors after the deploy.",
["billing", "technical", "sales", "other"]),
("A laptop was stolen from a coworking space.",
["low", "medium", "high", "critical"]),
]
flips = total = 0
for state, opts in cases:
q = {"name": "q", "type": "choice", "instructions": "Pick the best option.",
"options": opts, "answer": opts[0]}
base = opts[_predict(model, packer, state, q, device).argmax().item()]
for k in range(1, rotations):
rot = opts[k:] + opts[:k]
rq = dict(q, options=rot)
pick = rot[_predict(model, packer, state, rq, device).argmax().item()]
total += 1
flips += int(pick != base)
print(f"order: {flips}/{total} argmax changes under {rotations} rotations "
f"(flip rate {flips / total if total else 0:.3f})")
return flips / max(total, 1)
if __name__ == "__main__":
p = argparse.ArgumentParser()
p.add_argument("--checkpoint", required=True)
p.add_argument("--device", default="cpu")
p.add_argument("--rotations", type=int, default=6)
args = p.parse_args()
model = load_model(args.checkpoint, args.device)
packer = HFTokenPacker(model.cfg["backbone"])
print(f"probes on {args.checkpoint} ({args.device})\n")
isolation(model, packer, args.device)
print()
forgery(model, packer, args.device)
print()
order(model, packer, args.device, rotations=args.rotations)