botp
/

Solomon / src /solomon /readout.py
orz99's picture ArcherHume's picture
Duplicate from DoccyHealth/Solomon
1d2de8a
Raw History Blame Contribute Delete
14.1 kB
"""Contract v3 readouts: branch jobs for every answer type, predictions from letter logits, and metrics.
Job ids are '<row id>|<kind>[|<detail>]':
b / bp four-state Boolean question / its paraphrase
E|j L|j t|j entity j, label j, threshold j (four-state)
R|<perm> listwise with reserved options, caller options in order <perm> (digits; identity = caller order)
S|<perm> 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))}