"""Contract v3 readouts: branch jobs for every answer type, predictions from letter logits, and metrics. Job ids are '|[|]': b / bp four-state Boolean question / its paraphrase E|j L|j t|j entity j, label j, threshold j (four-state) R| listwise with reserved options, caller options in order (digits; identity = caller order) S| listwise without reserved options; suf = the separate sufficiency branch P|i per-option four-state branch for option i """ import random import zlib import numpy as np from solomon.engine_contract import CONFLICTING, NOT_STATED, four_state_block, label_block, listwise_block, option_block, sufficiency_block def softmax(x): x = np.asarray(x, np.float64) e = np.exp(x - x.max()) return e / e.sum() def perm_key(perm): return ''.join(str(i) for i in perm) def order_set(row_id, n, extra=3): """Caller order, every other cyclic rotation, and `extra` seeded random permutations (deduplicated).""" orders = [tuple((i + r) % n for i in range(n)) for r in range(n)] rng = random.Random(zlib.crc32(row_id.encode())) tries = 0 while len(orders) < n + extra and tries < 50: p = list(range(n)) rng.shuffle(p) if tuple(p) not in orders: orders.append(tuple(p)) tries += 1 return orders def rotations_of(perm): n = len(perm) return [tuple(perm[(i + r) % n] for i in range(n)) for r in range(n)] def jobs_for(rows, designs, docs=None): """designs: subset of {'four', 'para', 'R', 'Rrot', 'Rrotfull', 'S', 'Srot', 'suf', 'P'}. docs maps document text -> parts.""" out = [] for r in rows: inp = r['input'] doc = inp['document']['text'] if docs is None else docs[inp['document']['text']] add = lambda kind, blk: out.append({'id': f"{r['id']}|{kind}", 'doc': doc, 'block': blk[0], 'n': blk[1]}) task = r['task'] if task == 'boolean' and 'four' in designs: add('b', four_state_block(inp['question'])) if 'para' in designs and inp.get('paraphrase'): add('bp', four_state_block(inp['paraphrase'])) elif task == 'entity' and 'four' in designs: for j, name in enumerate(inp['entities']): add(f'E|{j}', four_state_block(inp['template'].format(entity=name))) elif task == 'multilabel' and 'four' in designs: for j, label in enumerate(inp['labels']): add(f'L|{j}', label_block(inp['question'], label)) elif task in ('single', 'ordered'): opts = inp['options'] ident = tuple(range(len(opts))) for d, reserved in (('R', True), ('S', False)): if not {d, d + 'rot', d + 'rotfull'} & set(designs): continue orders = [ident] if task == 'single' and (d + 'rot' in designs or d + 'rotfull' in designs): orders = order_set(r['id'], len(opts)) if d + 'rotfull' in designs: orders = list(dict.fromkeys(o for p in orders for o in rotations_of(p))) for perm in orders: add(f'{d}|{perm_key(perm)}', listwise_block(inp['question'], [opts[i] for i in perm], ordered=task == 'ordered', reserved=reserved)) if 'suf' in designs: add('suf', sufficiency_block(inp['question'], opts)) if 'P' in designs and task == 'single': for i, o in enumerate(opts): add(f'P|{i}', option_block(inp['question'], o)) if task == 'ordered' and 'four' in designs: for j, t in enumerate(inp.get('thresholds', [])): add(f't|{j}', four_state_block(t['question'])) return out # ---------------------------------------------------------------- predictions def listwise_logprobs(scores, row, design, perm, prior=None): """Log-probabilities over [option 0..n-1 in CALLER indexing] + reserved (if design R), from the branch scored in order `perm`.""" n = len(row['input']['options']) logits = np.asarray(scores[f"{row['id']}|{design}|{perm_key(perm)}"]['letter_logits'], np.float64) if prior is not None: logits = logits - prior[n if design == 'S' else n + 2] lp = logits - np.logaddexp.reduce(logits) out = np.empty_like(lp) for pos, i in enumerate(perm): out[i] = lp[pos] out[n:] = lp[n:] return out def decode(lp, n): k = int(np.argmax(lp)) return k if k < n else (NOT_STATED if k == n else CONFLICTING) def predict_listwise(scores, row, design='R', perm=None, average=None, prior=None): """average: None (one pass), or k = number of cyclic rotations of `perm` to average in log space ('all' = every rotation).""" n = len(row['input']['options']) perm = tuple(range(n)) if perm is None else perm if average is None: lp = listwise_logprobs(scores, row, design, perm, prior) else: rots = rotations_of(perm) k = n if average == 'all' else min(average, n) picks = [rots[round(j * n / k) % n] for j in range(k)] lp = np.mean([listwise_logprobs(scores, row, design, p, prior) for p in dict.fromkeys(picks)], axis=0) if design == 'R': return decode(lp, n) suf = int(np.argmax(scores[f"{row['id']}|suf"]['letter_logits'])) return int(np.argmax(lp[:n])) if suf == 0 else (NOT_STATED if suf == 1 else CONFLICTING) def predict_per_option(scores, row): n = len(row['input']['options']) p = np.stack([softmax(scores[f"{row['id']}|P|{i}"]['letter_logits']) for i in range(n)]) yes, neither = p[:, 0] + p[:, 3], p[:, 2] if (yes > 0.5).sum() >= 2: return CONFLICTING if not (yes > neither).any(): return NOT_STATED return int(np.argmax(p[:, 0])) def position_prior(scores, rows, design='R'): """PriDe-style prior over letter positions, per option count: mean log-probability of each position over all cyclic rotations of the estimation rows (every option visits every position, so content averages out).""" acc = {} for r in rows: if r['task'] != 'single': continue n = len(r['input']['options']) for perm in rotations_of(tuple(range(n))): key = f"{r['id']}|{design}|{perm_key(perm)}" if key not in scores: break l = np.asarray(scores[key]['letter_logits'], np.float64) acc.setdefault(len(l), []).append(l - np.logaddexp.reduce(l)) prior = {} for width, vals in acc.items(): m = np.mean(vals, axis=0) n = width - 2 if design == 'R' else width prior[width] = np.zeros(width) # reserved positions are never rotated: left unadjusted prior[width][:n] = m[:n] - m[:n].mean() return prior def four_state_pred(scores, key): return int(np.argmax(scores[key]['letter_logits'][:4])) def derived_threshold(scores, row, k, design='R'): """P(level >= k | some level) from the listwise distribution: monotone in k by construction.""" n = len(row['input']['options']) p = np.exp(listwise_logprobs(scores, row, design, tuple(range(n))))[:n] return float(p[k:].sum() / p.sum()) # ---------------------------------------------------------------- metrics def bootstrap(correct_a, correct_b, families, n=2000, seed=4): rng = np.random.default_rng(seed) correct_a, correct_b, families = np.asarray(correct_a, float), np.asarray(correct_b, float), np.asarray(families) fams = np.unique(families) idx = {f: np.flatnonzero(families == f) for f in fams} diffs = [] for _ in range(n): take = np.concatenate([idx[f] for f in rng.choice(fams, len(fams), replace=True)]) diffs.append(correct_a[take].mean() - correct_b[take].mean()) lo, hi = np.percentile(diffs, [2.5, 97.5]) return {'difference': float(correct_a.mean() - correct_b.mean()), 'ci95': [float(lo), float(hi)]} def interval(correct, families, n=2000, seed=4): rng = np.random.default_rng(seed) correct, families = np.asarray(correct, float), np.asarray(families) fams = np.unique(families) idx = {f: np.flatnonzero(families == f) for f in fams} vals = [correct[np.concatenate([idx[f] for f in rng.choice(fams, len(fams), replace=True)])].mean() for _ in range(n)] lo, hi = np.percentile(vals, [2.5, 97.5]) return [float(lo), float(hi)] def decisions(rows, scores, design='R', average=None, prior=None, per_option=False): """One record per scored decision: {'row', 'task', 'unit', 'family', 'generator', 'panel', 'gold', 'pred', 'correct', ...}.""" out = [] for r in rows: base = {'row': r['id'], 'task': r['task'], 'family': r['family_id'], 'generator': r['generator'], 'panel': r['panel'], 'mechanism': r.get('mechanism')} if r['task'] == 'boolean': pred = four_state_pred(scores, f"{r['id']}|b") out.append({**base, 'unit': 'b', 'gold': r['target'], 'pred': pred, 'correct': pred == r['target']}) elif r['task'] == 'entity': for j, gold in enumerate(r['targets']): pred = four_state_pred(scores, f"{r['id']}|E|{j}") out.append({**base, 'unit': f'E|{j}', 'gold': gold, 'pred': pred, 'correct': pred == gold, 'absent': r['absent'][j]}) elif r['task'] == 'multilabel': for j, gold in enumerate(r['targets']): pred = four_state_pred(scores, f"{r['id']}|L|{j}") out.append({**base, 'unit': f'L|{j}', 'gold': gold, 'pred': pred, 'correct': pred == gold}) else: pred = predict_per_option(scores, r) if (per_option and r['task'] == 'single') else predict_listwise( scores, r, design, average=average if r['task'] == 'single' else None, prior=prior if r['task'] == 'single' else None) out.append({**base, 'unit': 'choice', 'gold': r['target'], 'pred': pred, 'correct': pred == r['target'], 'reserved_gold': not isinstance(r['target'], int), 'n_options': len(r['input']['options'])}) return out def summarise(decs): out = {} for task in sorted({d['task'] for d in decs}): sel = [d for d in decs if d['task'] == task] ok = [d['correct'] for d in sel] entry = {'correct': int(sum(ok)), 'n': len(sel), 'accuracy': float(np.mean(ok)), 'ci95': interval(ok, [d['family'] for d in sel]), 'by_generator': {}, 'by_panel': {}} for key in ('generator', 'panel'): for v in sorted({d[key] for d in sel}): s = [d['correct'] for d in sel if d[key] == v] entry['by_' + key][v] = f'{int(sum(s))}/{len(s)}' if task == 'boolean': entry['by_mechanism'] = {m: f"{sum(d['correct'] for d in sel if d['mechanism'] == m)}/{sum(d['mechanism'] == m for d in sel)}" for m in sorted({d['mechanism'] for d in sel})} if task in ('single', 'ordered'): for name, flag in (('reserved_gold', True), ('index_gold', False)): s = [d['correct'] for d in sel if d['reserved_gold'] == flag] entry[name] = f'{int(sum(s))}/{len(s)}' if task in ('multilabel', 'entity'): groups = {} for d in sel: groups.setdefault(d['row'], []).append(d['correct']) entry['all_units_correct'] = f'{sum(all(v) for v in groups.values())}/{len(groups)}' if task == 'entity': s = [d['correct'] for d in sel if d['absent']] entry['absent_entities'] = f'{int(sum(s))}/{len(s)}' out[task] = entry return out def flip_rate(rows, scores, design='R', average=None, prior=None): """Share of single-choice rows whose top answer (as an option identity) changes under any order in the order set.""" flips, n = [], 0 for r in rows: if r['task'] != 'single': continue opts = len(r['input']['options']) base = predict_listwise(scores, r, design, None, average, prior) preds = [predict_listwise(scores, r, design, perm, average, prior) for perm in order_set(r['id'], opts)[1:]] flips.append(any(p != base for p in preds)) n += 1 return {'rows': n, 'flipped': int(sum(flips)), 'rate': float(np.mean(flips)) if flips else None} def threshold_consistency(rows, scores, design='R'): agree, total, derived_ok, asked_ok, decided = 0, 0, 0, 0, 0 for r in rows: if r['task'] != 'ordered': continue for j, t in enumerate(r['input']['thresholds']): asked = four_state_pred(scores, f"{r['id']}|t|{j}") gold = r['threshold_targets'][j] asked_ok += asked == gold total += 1 if gold in (0, 1): d = 0 if derived_threshold(scores, r, t['k'], design) >= 0.5 else 1 decided += 1 derived_ok += d == gold agree += d == asked return {'threshold_questions': total, 'asked_correct': int(asked_ok), 'decided_gold': decided, 'derived_correct': int(derived_ok), 'derived_agrees_with_asked': int(agree)} def paraphrase_consistency(rows, scores): same = [four_state_pred(scores, f"{r['id']}|b") == four_state_pred(scores, f"{r['id']}|bp") for r in rows if r['task'] == 'boolean' and f"{r['id']}|bp" in scores] return {'pairs': len(same), 'same_answer': int(sum(same))} def negation_consistency(rows, scores): """Positive and negative query of one subject: the four-state answers must be swaps of each other.""" swap = {0: 1, 1: 0, 2: 2, 3: 3} pairs = {} for r in rows: if r['task'] == 'boolean' and 'subject' in r: pairs.setdefault((r['family_id'], r['subject']), {})[r['negative_query']] = four_state_pred(scores, f"{r['id']}|b") ok = [swap[v[0]] == v[1] for v in pairs.values() if set(v) == {0, 1}] return {'pairs': len(ok), 'consistent': int(sum(ok))}