-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsvr_mad.py
More file actions
455 lines (400 loc) · 18.7 KB
/
Copy pathsvr_mad.py
File metadata and controls
455 lines (400 loc) · 18.7 KB
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
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
"""SVR-MAD: pick receiver by argmax Corr (initialized from --prior-signal),
fire up to --P d=1 probes against its top disagreeing peers excluding
already-debated (ranked by --peer-signal), and update Corr by SVR =
(retain - flip) / D. Commit the receiver's answer when
Corr >= --acceptance-thres and D >= --min-peers-to-accept; otherwise
fall back to --tiebreak."""
import argparse
import asyncio
import logging
import sys
from collections import Counter
from dataclasses import dataclass, field
from datetime import datetime, timezone
from pathlib import Path
from tqdm import tqdm
from api_client import BedrockClient
from credentials import get_api_key
from data_loader import load_predebate
from debate_io import already_done, read_predebate_config, write_debate_result
from math_canonicalizer import MathCanonicalizer, canon
from peer_msg import build_peer_block, peer_message, pick_header
from prompt_template import QUERY_TEMPLATE_MATH, QUERY_TEMPLATE_MC
TEMPLATES = {"math": QUERY_TEMPLATE_MATH, "mc": QUERY_TEMPLATE_MC}
log = logging.getLogger("svr_mad")
def parse_args():
p = argparse.ArgumentParser(description="SVR-MAD: probe + commit + fallback.")
p.add_argument("--predebate-dir", required=True, dest="predebate_dir")
p.add_argument("--P", type=int, default=2, help="Probes per agent (default 2).")
p.add_argument("--prior-signal", choices=["perplexity", "min_logprob", "confidence"],
default="perplexity", dest="prior_signal",
help="Agent iteration order. perplexity asc (low=confident first); "
"min_logprob desc (high=confident); confidence desc.")
p.add_argument("--peer-signal", choices=["perplexity", "min_logprob", "confidence"],
default="perplexity", dest="peer_signal",
help="Within-agent cross-peer rank. Most-confident-first: "
"perplexity asc, min_logprob desc, confidence desc.")
p.add_argument("--budget", type=int, default=None,
help="Max LLM calls per question. Overrides --budget-rule.")
p.add_argument("--budget-rule", choices=["fixed", "clusters"], default="clusters",
dest="budget_rule",
help="If --budget is not set: fixed=n*P (uncapped); "
"clusters=P*(n_distinct_answers + max_cluster).")
p.add_argument("--tiebreak", choices=["round0_mv", "layered_mv"], default="layered_mv",
help="Fallback when no agent commits. Default layered_mv.")
p.add_argument("--acceptance-thres", type=float, default=1.0, dest="acceptance_thres",
help="Acceptance threshold tau: commit when Corr[A_r] >= tau (default 1.0).")
p.add_argument("--min-peers-to-accept", type=int, default=None, dest="min_peers_to_accept",
help="Minimum probes C required before commit. Defaults to --P.")
p.add_argument("--include-reasoning", action="store_true", dest="include_reasoning")
p.add_argument("--only-non-unanimous", action="store_true", dest="only_non_unanimous")
p.add_argument("--convert-math-answers", action="store_true", dest="convert_math_answers",
help="Math only: cluster boxed_answers by math equivalence so the "
"cross filter, commit check, and MV fallback all run on canonical "
"strings. Every snapshot gets parsed_math_answer.")
p.add_argument("--peer-header",
choices=["math_full", "math_full_strong", "mc_full", "mc_full_strong"],
default=None)
p.add_argument("--temperature", type=float, default=None)
p.add_argument("--top-p", type=float, default=None, dest="top_p")
p.add_argument("--top-k", type=int, default=None, dest="top_k")
p.add_argument("--max-tokens", type=int, default=None, dest="max_tokens")
p.add_argument("--reasoning-effort", choices=["low", "medium", "high"],
default=None, dest="reasoning_effort")
p.add_argument("--seed-base", type=int, default=None, dest="seed_base",
help="Per-call seed = seed_base + receiver*100 + peer; omit for unseeded.")
p.add_argument("--concurrent-queries", type=int, default=3, dest="concurrent_queries")
p.add_argument("--concurrent-requests", type=int, default=20, dest="concurrent_requests")
p.add_argument("--limit", type=int, default=None)
p.add_argument("--output-dir", default=None, dest="output_dir")
p.add_argument("-v", "--verbose", action="store_true")
args = p.parse_args()
if args.min_peers_to_accept is None:
args.min_peers_to_accept = args.P
return args
def setup_logging(out_dir, verbose):
log.setLevel(logging.DEBUG)
fmt = logging.Formatter("%(asctime)s %(levelname)s %(message)s", datefmt="%H:%M:%S")
sh = logging.StreamHandler(sys.stdout)
sh.setLevel(logging.DEBUG if verbose else logging.INFO)
sh.setFormatter(fmt)
fh = logging.FileHandler(out_dir / "svr_mad_run.log")
fh.setLevel(logging.DEBUG)
fh.setFormatter(fmt)
log.handlers = [sh, fh]
def resolve_config(args, pre):
pick = lambda override, fallback: override if override is not None else fallback
return {
"model": pre["model"],
"model_api_id": pre["model_api_id"],
"task": pre["task"],
"dataset": pre["dataset"],
"n": pre["n"],
"temperature": pick(args.temperature, pre["temperature"]),
"top_p": pick(args.top_p, pre["top_p"]),
"top_k": pick(args.top_k, pre["top_k"]),
"max_tokens": pick(args.max_tokens, pre["max_tokens"]),
"reasoning_effort": pick(args.reasoning_effort, pre.get("reasoning_effort")),
}
def rank_agents_by_prior(self_rounds, signal):
"""Most-confident-first order of agent ids."""
n = len(self_rounds)
if signal == "perplexity":
return sorted(range(n), key=lambda i:
(self_rounds[i].get("logprob_stats") or {}).get("perplexity") or float("inf"))
if signal == "min_logprob":
return sorted(range(n), key=lambda i: -(
(self_rounds[i].get("logprob_stats") or {}).get("min_logprob") or float("-inf")))
if signal == "confidence":
return sorted(range(n), key=lambda i: -(self_rounds[i].get("confidence") or 0))
raise ValueError(f"unknown prior signal: {signal}")
def rank_cross(answers, self_rounds, n, i, signal, debated=None):
"""Most-confident-first ordering of cross peers, excluding already-debated."""
own = answers[i]
skip = debated or set()
cross = [j for j in range(n) if j != i and answers[j] != own and j not in skip]
if signal == "perplexity":
return sorted(cross, key=lambda j:
(self_rounds[j].get("logprob_stats") or {}).get("perplexity") or float("inf"))
if signal == "min_logprob":
return sorted(cross, key=lambda j: -(
(self_rounds[j].get("logprob_stats") or {}).get("min_logprob") or float("-inf")))
return sorted(cross, key=lambda j: -(self_rounds[j].get("confidence") or 0))
def _mv(items, prior=None):
"""Majority vote with prior as tiebreak; lex-smallest if no prior fits."""
items = [x for x in items if x is not None]
if not items:
return prior
counts = Counter(items)
top = max(counts.values())
cands = [a for a, c in counts.items() if c == top]
if len(cands) == 1:
return cands[0]
if prior is not None and prior in cands:
return prior
return sorted(cands)[0]
def apply_tiebreak(method, answers, probes_per_agent, n, answer_key):
"""`answer_key(snap)` returns the comparison key (canonical or raw) used
consistently for `answers` and probe snapshots."""
if method == "round0_mv":
return _mv(answers)
if method == "layered_mv":
per_agent = []
for i in range(n):
probe_answers = [answer_key(pr["post_debate_snapshot"])
for pr in probes_per_agent.get(i, [])]
per_agent.append(_mv(probe_answers, prior=answers[i]))
return _mv(per_agent, prior=_mv(answers))
raise ValueError(f"unknown tiebreak: {method}")
# --- Per-agent state --------------------------------------------------------
# D, r, c: probes observed, retentions, changes.
# debated: peer ids already probed for this receiver. Agent leaves the
# candidate pool when |debated| >= S (per-receiver quota) or no remaining
# disagreeing peers exist outside debated.
#
# Corr[A] = SVR(D, r, c) = (r - c) / D if D > 0
# = Prior(A) = -prior_rank otherwise
@dataclass
class AgentState:
agent_id: int
prior_rank: int
D: int = 0
r: int = 0
c: int = 0
debated: set = field(default_factory=set)
def corr_score(state):
if state.D > 0:
return (state.r - state.c) / state.D
return -state.prior_rank
def pick_receiver(states, answers, n, S):
"""A_r = argmax_A Corr[A] over agents with debate quota and disagreeing peers left."""
candidates = []
for s in states:
if len(s.debated) >= S:
continue
own = answers[s.agent_id]
if not any(j != s.agent_id and answers[j] != own and j not in s.debated
for j in range(n)):
continue
candidates.append(s)
if not candidates:
return None
return max(candidates, key=corr_score)
async def run_one_question(client, query, predebate, cfg, args, header):
n = cfg["n"]
template = TEMPLATES[cfg["task"]]
base_msg = template.format(Question=query["question"])
self_rounds = [ag["self_round"] for ag in predebate["agents"][:n]]
gt = query["answer"]
do_convert = args.convert_math_answers and cfg["task"] == "math"
if do_convert:
canonicalizer = MathCanonicalizer(gt, seed=[sr.get("boxed_answer") for sr in self_rounds])
parsed_gt = canon(gt)
else:
canonicalizer = None
parsed_gt = gt
def answer_key(snap_like):
ans = snap_like.get("boxed_answer")
if do_convert and ans is not None:
return canonicalizer.mapping.get(ans, ans)
return ans
answers = [answer_key(sr) for sr in self_rounds]
# Corr[A_i] = Prior(A_i).
iter_order = rank_agents_by_prior(self_rounds, args.prior_signal)
prior_rank_of = {a: r for r, a in enumerate(iter_order)}
states = [AgentState(agent_id=i, prior_rank=prior_rank_of[i]) for i in range(n)]
# Snapshot of top-S Disagree per receiver at t=0 (for output JSON only;
# the loop re-derives peers dynamically against ar.debated each iteration).
S = args.P
probe_peers_planned = {
i: rank_cross(answers, self_rounds, n, i, args.peer_signal)[:S]
for i in range(n)
}
# B_max = S * (k + m).
if args.budget is not None:
budget = args.budget
elif args.budget_rule == "clusters":
counts = Counter(a for a in answers if a is not None)
n_clusters = len(counts)
max_cluster = max(counts.values()) if counts else 0
budget = S * (n_clusters + max_cluster)
else: # fixed
budget = n * S
# Force budget to a multiple of S so partial-batch tails can't fire.
budget = (budget // S) * S
probes_per_agent = {i: [] for i in range(n)}
comms = 0
committed_agent = None
committed_answer = None
async def fire_probe(receiver, peer):
peer_items = [(peer, peer_message(self_rounds[peer],
include_reasoning=args.include_reasoning))]
peer_block = build_peer_block(peer_items, header)
msgs = [
{"role": "user", "content": base_msg},
{"role": "assistant", "content": self_rounds[receiver].get("content") or ""},
{"role": "user", "content": peer_block},
]
seed = None if args.seed_base is None else args.seed_base + receiver * 100 + peer
snap = await client.chat(
model=cfg["model_api_id"],
messages=msgs,
temperature=cfg["temperature"], top_p=cfg["top_p"], top_k=cfg["top_k"],
max_tokens=cfg["max_tokens"], seed=seed,
reasoning_effort=cfg["reasoning_effort"],
)
if do_convert:
snap["parsed_math_answer"] = canonicalizer.add(snap.get("boxed_answer"))
return snap
while comms < budget:
# A_r = argmax_A Corr[A].
ar = pick_receiver(states, answers, n, S)
if ar is None:
break
A = ar.agent_id
# P = top-(S - |Debated[r]|) of (Disagree \ Debated[r]) by peer_signal.
avail = rank_cross(answers, self_rounds, n, A, args.peer_signal, ar.debated)
remaining = budget - comms
cap = S - len(ar.debated)
peers_to_fire = avail[: min(cap, remaining)]
if not peers_to_fire:
continue
# Probe peers; update Corr[A_r] = SVR(D, r, c).
snaps = await asyncio.gather(*[fire_probe(A, p) for p in peers_to_fire])
comms += len(snaps)
own = answers[A]
for p, snap in zip(peers_to_fire, snaps):
probes_per_agent[A].append({"peer_id": p, "post_debate_snapshot": snap})
ar.D += 1
if answer_key(snap) == own:
ar.r += 1
else:
ar.c += 1
ar.debated.add(p)
# Commit if Corr[A_r] >= tau and D >= C.
if ar.D >= args.min_peers_to_accept and corr_score(ar) >= args.acceptance_thres:
committed_agent = A
committed_answer = own
break
# Refresh answers from the canonicalizer's current mapping before
# tiebreaking; cluster lex-mins may have shifted as probes added members.
if do_convert:
answers_live = [answer_key(sr) for sr in self_rounds]
else:
answers_live = answers
fallback_used = committed_answer is None
final_answer = committed_answer
if fallback_used:
final_answer = apply_tiebreak(args.tiebreak, answers_live, probes_per_agent, n, answer_key)
elif do_convert:
final_answer = answers_live[committed_agent]
# Re-stamp every probe snapshot's parsed_math_answer from the final mapping.
if do_convert:
for ps in probes_per_agent.values():
for pr in ps:
snap = pr["post_debate_snapshot"]
ans = snap.get("boxed_answer")
if ans is not None:
snap["parsed_math_answer"] = canonicalizer.mapping.get(ans)
total_in = sum(pr["post_debate_snapshot"].get("input_tokens", 0)
for ps in probes_per_agent.values() for pr in ps)
total_out = sum(pr["post_debate_snapshot"].get("output_tokens", 0)
for ps in probes_per_agent.values() for pr in ps)
return {
"query_id": query["id"],
"question": query["question"],
"ground_truth_answer": gt,
"parsed_ground_truth_answer": parsed_gt,
"extra": query["extra"],
"algorithm": "svr_mad",
"predebate_dir": args.predebate_dir,
"config": {
**cfg,
"P": args.P,
"prior_signal": args.prior_signal,
"peer_signal": args.peer_signal,
"budget": args.budget,
"tiebreak": args.tiebreak,
"acceptance_thres": args.acceptance_thres,
"min_peers_to_accept": args.min_peers_to_accept,
"include_reasoning": args.include_reasoning,
"only_non_unanimous": args.only_non_unanimous,
"peer_header_variant": args.peer_header,
"seed_base": args.seed_base,
"convert_math_answers": args.convert_math_answers,
},
"timestamp": datetime.now(timezone.utc).isoformat(),
"total_input_tokens": total_in,
"total_output_tokens": total_out,
"final_answer": final_answer,
"committed_agent": committed_agent,
"fallback_used": fallback_used,
"comms_used": comms,
"comms_budget": budget,
"agent_iter_order": iter_order,
"agents": [
{
"agent_id": i,
"predebate_boxed_answer": answers[i],
"probe_peers_planned": probe_peers_planned[i],
"probes": probes_per_agent[i],
}
for i in range(n)
],
}
async def main():
args = parse_args()
pre_cfg = read_predebate_config(args.predebate_dir)
cfg = resolve_config(args, pre_cfg)
if args.output_dir:
out_dir = Path(args.output_dir)
else:
parent = Path(args.predebate_dir).resolve().parent
out_dir = parent / f"svr_mad_n{cfg['n']}_P{args.P}_{args.prior_signal}"
out_dir.mkdir(parents=True, exist_ok=True)
setup_logging(out_dir, args.verbose)
log.info("Output dir: %s", out_dir)
log.info("Resolved config: %s", cfg)
log.info("P=%d prior_signal=%s peer_signal=%s budget=%s tiebreak=%s "
"acceptance_thres=%g min_peers_to_accept=%d",
args.P, args.prior_signal, args.peer_signal,
args.budget if args.budget is not None else "uncapped", args.tiebreak,
args.acceptance_thres, args.min_peers_to_accept)
items = load_predebate(args.predebate_dir, only_non_unanimous=args.only_non_unanimous)
if args.limit:
items = items[: args.limit]
done = already_done(out_dir)
pending = [it for it in items if it[0]["id"] not in done]
log.info("Total: %d Done: %d Pending: %d", len(items), len(done), len(pending))
if not pending:
log.info("Nothing to do; all queries already processed.")
return
header_variant = args.peer_header or ("math_full" if cfg["task"] == "math" else "mc_full")
header = pick_header(cfg["task"], header_variant)
log.info("Peer header variant: %s", header_variant)
api_key = get_api_key()
outer_sem = asyncio.Semaphore(args.concurrent_queries)
async with BedrockClient(api_key, concurrent_requests=args.concurrent_requests) as client:
pbar = tqdm(total=len(pending), unit="q")
async def run(item):
query, predebate = item
async with outer_sem:
try:
result = await run_one_question(client, query, predebate, cfg, args, header)
write_debate_result(out_dir, query["id"], result)
pbar.set_postfix(
tok=f"{result['total_input_tokens']}in/{result['total_output_tokens']}out",
comms=result["comms_used"],
commit=str(result["committed_agent"] if not result["fallback_used"] else "fb"),
)
log.debug("done %s", query["id"])
except Exception as e:
pbar.write(f"ERROR [{query['id']}]: {type(e).__name__}: {e}")
log.exception("Failed on %s", query["id"])
finally:
pbar.update(1)
await asyncio.gather(*[run(item) for item in pending])
pbar.close()
log.info("Run complete.")
if __name__ == "__main__":
asyncio.run(main())