opensysone / source /selection.py
andyshu's picture
Back up verified OpenSysOne training snapshot and pinned source
2d5c26a verified
Raw History Blame
5.34 kB
"""Frozen validation-only checkpoint selection; no deployment calibration reuse.
Fit one scalar temperature on three validation folds and score the fourth. Source
groups never cross folds. Both fitting and final selection weight task families
equally. The final serving temperature is fitted later on reserved calibration.
"""
from collections import defaultdict
import math
import random
import statistics
SELECTION_METRIC = 'crossfit_temperature_nll_v1'
FOLDS = 4
SEED = 431
TEMPERATURES = tuple(10.0 ** (-1.0 + 2.3 * index / 100) for index in range(101))
SELECTION_POLICY = {'id': SELECTION_METRIC, 'folds': FOLDS, 'seed': SEED,
'grouping': 'source group, nested within task family; groups shuffled within sorted families and assigned round-robin',
'temperature_grid': {'count': 101, 'log10_min': -1.0, 'log10_max': 1.3, 'arithmetic': 'Python float64'},
'temperature_fit': 'minimum macro-family NLL on the other three folds; smallest temperature wins ties',
'score': 'macro-family mean of all held-out decision NLLs',
'deployment_temperature': 'fit afresh on reserved calibration only after model selection'}
def _nll(logits, target, temperature):
maximum = max(logits)
scaled = [(value - maximum) / temperature for value in logits]
return math.log(math.fsum(math.exp(value) for value in scaled)) - scaled[target]
def validation_selection(predictions):
"""Return a deterministic, calibration-aware selection score from raw logits."""
if not predictions:
raise ValueError('Selection needs validation predictions')
rows = sorted(predictions, key=lambda row: row['id'])
seen, group_family = set(), {}
families = defaultdict(list)
groups = defaultdict(lambda: defaultdict(list))
curves, raw, correct = [], [], []
for index,row in enumerate(rows):
identifier, group, family = row['id'], row['group'], row['family']
if not all(isinstance(value,str) and value for value in (identifier,group,family)):
raise ValueError('Validation identity, group and family must be nonempty strings')
if identifier in seen:
raise ValueError('Validation IDs must be unique')
seen.add(identifier)
if group in group_family and group_family[group] != family:
raise ValueError('Source groups must be nested within task families')
group_family[group] = family
logits, target = row['logits'], row['target']
if len(logits) < 2 or any(not isinstance(value,(float,int)) or not math.isfinite(value) for value in logits):
raise ValueError('Selection needs finite logits and at least two choices')
if not isinstance(target,int) or isinstance(target,bool) or not 0 <= target < len(logits):
raise ValueError('Invalid target index')
families[family].append(index)
groups[family][group].append(index)
curves.append([_nll(logits,target,temperature) for temperature in TEMPERATURES])
raw.append(_nll(logits,target,1.0))
correct.append(float(max(range(len(logits)),key=logits.__getitem__) == target))
assignments = [None] * len(rows)
rng = random.Random(SEED)
for family in sorted(groups):
keys = sorted(groups[family])
if len(keys) < FOLDS:
raise ValueError('Each family needs at least four independent source groups')
rng.shuffle(keys)
for position,group in enumerate(keys):
for index in groups[family][group]:
assignments[index] = position % FOLDS
heldout_losses = [None] * len(rows)
fold_results = []
for fold in range(FOLDS):
training = {family:[i for i in indices if assignments[i] != fold]
for family,indices in families.items()}
objective = [statistics.mean(math.fsum(curves[i][temperature] for i in indices)/len(indices)
for indices in training.values()) for temperature in range(len(TEMPERATURES))]
best = min(range(len(TEMPERATURES)),key=objective.__getitem__)
held = [i for i,assignment in enumerate(assignments) if assignment == fold]
for index in held:
heldout_losses[index] = curves[index][best]
held_groups = {rows[index]['group'] for index in held}
fold_results.append({'fold':fold,'temperature_index':best,'temperature':TEMPERATURES[best],
'training_macro_nll':objective[best],'heldout_ids':[rows[index]['id'] for index in held],
'heldout_group_count':len(held_groups),'training_group_count':len(group_family)-len(held_groups)})
per_family = {family:{'count':len(indices),
'score':math.fsum(heldout_losses[i] for i in indices)/len(indices),
'raw_macro_nll':math.fsum(raw[i] for i in indices)/len(indices),
'accuracy':math.fsum(correct[i] for i in indices)/len(indices)}
for family,indices in sorted(families.items())}
return {'metric':SELECTION_METRIC,'score':statistics.mean(value['score'] for value in per_family.values()),
'raw_macro_nll':statistics.mean(value['raw_macro_nll'] for value in per_family.values()),
'accuracy':math.fsum(correct)/len(correct),'folds':fold_results,'per_family':per_family,
'policy':{**SELECTION_POLICY,'temperature_grid':dict(SELECTION_POLICY['temperature_grid'])}}