"""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'])}}