coco-comparator-bank / tests /test_rules.py
CharlesCNorton
One coordinate comparison per COCO category on EUPE-ViT-B features
8f61e02
Raw
History Blame Contribute Delete
2.29 kB
"""The rules are well formed and say the same thing everywhere they appear."""
import pytest
from conftest import load
NAMES = sorted(load('rules.json')['rules'])
@pytest.mark.parametrize('name', NAMES)
def test_the_two_dims_differ(name, rules):
r = rules['rules'][name]
assert r['pos_dim'] != r['neg_dim'], f'{name} compares a channel with itself'
@pytest.mark.parametrize('name', NAMES)
def test_dims_are_in_range(name, rules):
r = rules['rules'][name]
for d in (r['pos_dim'], r['neg_dim']):
assert 0 <= d < 768
@pytest.mark.parametrize('name', NAMES)
def test_category_ids_are_coco(name, rules):
assert 1 <= rules['rules'][name]['cat_id'] <= 90
def test_category_ids_are_unique(rules):
ids = [r['cat_id'] for r in rules['rules'].values()]
assert len(set(ids)) == len(ids)
@pytest.mark.parametrize('name', NAMES)
def test_eval_carries_the_same_rule(name, rules, evaluation):
r, e = rules['rules'][name], evaluation['rules'][name]
assert (e['pos_dim'], e['neg_dim']) == (r['pos_dim'], r['neg_dim'])
assert e['cat_id'] == r['cat_id']
assert e['train_balanced_accuracy'] == r['train_balanced_accuracy']
@pytest.mark.parametrize('name', NAMES)
def test_every_rule_beats_chance(name, evaluation):
"""Balanced accuracy on a balanced subsample; 0.5 is chance for any rate."""
assert evaluation['rules'][name]['balanced_accuracy'] > 0.5
@pytest.mark.parametrize('name', NAMES)
def test_scoring_set_is_balanced_and_large_enough(name, evaluation):
e = evaluation['rules'][name]
assert e['n_balanced'] % 2 == 0
assert e['n_balanced'] >= 40
def test_distribution_matches_the_rules(evaluation):
accs = sorted(v['balanced_accuracy'] for v in evaluation['rules'].values())
d = evaluation['distribution']
assert d['n_categories'] == len(accs)
assert d['min'] == accs[0] and d['max'] == accs[-1]
for t, n in d['at_least'].items():
assert n == sum(a >= float(t) for a in accs)
def test_skipped_categories_are_recorded(rules):
"""Categories dropped for lack of validation images are named, not silent."""
skipped = rules['provenance']['skipped']
assert len(rules['rules']) + len(skipped) == 80
for s in skipped:
assert s['val_positives'] < rules['min_val_images']