Download source/tests/test_selection.py from andyshu/opensysone: direct link, hf CLI and curl.
- Browser
- Download file 4.09 kB
-
https://huggingface.co/andyshu/opensysone/resolve/f2d6f8daa15bd21c316e61249f45ac16cbb79d45/source/tests/test_selection.py
- Command line
-
hf download hf://andyshu/opensysone@f2d6f8daa15bd21c316e61249f45ac16cbb79d45/source/tests/test_selection.py
-
curl -L -o test_selection.py https://huggingface.co/andyshu/opensysone/resolve/f2d6f8daa15bd21c316e61249f45ac16cbb79d45/source/tests/test_selection.py
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() | |