opensysone / source /tests /test_selection.py
andyshu's picture
Back up verified OpenSysOne training snapshot and pinned source
2d5c26a verified
Raw History Blame
4.09 kB
import copy
import random
import unittest
from selection import SELECTION_METRIC, TEMPERATURES, validation_selection
def rows():
return [{'id':f'{family}-{group:02d}-{index}','family':family,'group':f'{family}-{group:02d}',
'target':int((group+index)%3==0),'logits':[2.0,0.0]}
for family in ('a','b') for group in range(12) for index in range(1+group%3)]
class SelectionTests(unittest.TestCase):
def test_deterministic_order_invariant_and_whole_groups_stay_together(self):
predictions=rows()
result=validation_selection(predictions)
random.Random(77).shuffle(predictions)
self.assertEqual(result,validation_selection(predictions))
self.assertEqual(result['metric'],SELECTION_METRIC)
by_id={row['id']:row for row in predictions}
assignments={}
for fold in result['folds']:
for identifier in fold['heldout_ids']:
group=by_id[identifier]['group']
self.assertIn(assignments.get(group,fold['fold']),(fold['fold'],))
assignments[group]=fold['fold']
self.assertEqual(len(assignments),24)
self.assertEqual(sum(len(f['heldout_ids']) for f in result['folds']),len(predictions))
def test_held_out_labels_never_fit_their_own_temperature(self):
predictions=rows()
original=validation_selection(predictions)
for fold in original['folds']:
held=set(fold['heldout_ids'])
changed=copy.deepcopy(predictions)
for row in changed:
if row['id'] in held:
row['target']=1-row['target']
result=validation_selection(changed)
self.assertEqual(result['folds'][fold['fold']]['temperature_index'],fold['temperature_index'])
self.assertEqual(result['folds'][fold['fold']]['training_macro_nll'],fold['training_macro_nll'])
def test_temperature_absorbs_a_grid_step_global_logit_scale(self):
predictions=rows()
original=validation_selection(predictions)
factor=TEMPERATURES[51]/TEMPERATURES[50]
changed=copy.deepcopy(predictions)
for row in changed:
row['logits']=[value*factor for value in row['logits']]
result=validation_selection(changed)
self.assertAlmostEqual(result['score'],original['score'],places=12)
for before,after in zip(original['folds'],result['folds']):
self.assertTrue(0 < before['temperature_index'] < 99)
self.assertEqual(after['temperature_index'],before['temperature_index']+1)
def test_macro_score_equal_weights_families_with_unequal_groups_and_rows(self):
predictions=rows()
predictions += [{**row,'id':row['id']+'-copy'} for row in predictions if row['family']=='a']
result=validation_selection(predictions)
self.assertNotEqual(result['per_family']['a']['count'],result['per_family']['b']['count'])
self.assertAlmostEqual(result['score'],sum(v['score'] for v in result['per_family'].values())/2)
# Duplicating every decision in one family must not increase its fit weight.
base=validation_selection(rows())
self.assertEqual([f['temperature_index'] for f in base['folds']],
[f['temperature_index'] for f in result['folds']])
self.assertAlmostEqual(base['score'],result['score'],places=12)
def test_rejects_cross_family_groups_duplicate_ids_and_nonfinite_logits(self):
original=rows()
cases=[]
changed=copy.deepcopy(original);changed[-1]['group']=changed[0]['group'];cases.append(changed)
changed=copy.deepcopy(original);changed[-1]['id']=changed[0]['id'];cases.append(changed)
changed=copy.deepcopy(original);changed[-1]['logits'][0]=float('nan');cases.append(changed)
for changed in cases:
with self.assertRaises(ValueError):validation_selection(changed)
with self.assertRaisesRegex(ValueError,'four independent'):
validation_selection(original[:2])
if __name__=='__main__':unittest.main()