"""The 40 classifier dims agree everywhere they are written down.""" from safetensors.torch import load_file from conftest import REPO, load N_POS = 20 N_NEG = 20 def test_baseline_counts(classifier): assert len(classifier['pos_dims']) == N_POS assert len(classifier['neg_dims']) == N_NEG assert not set(classifier['pos_dims']) & set(classifier['neg_dims']) assert classifier['fixed_parameters']['dim_indices'] == N_POS + N_NEG assert classifier['free_parameters'] == 1 def test_tight_fpr_extends_the_baseline(classifier, tight): assert tight['pos_dims'] == classifier['pos_dims'] assert tight['neg_dims_original'] == classifier['neg_dims'] assert tight['neg_dims'] == tight['neg_dims_original'] + tight['neg_dims_extra'] assert len(set(tight['neg_dims'])) == len(tight['neg_dims']) assert tight['fixed_parameters']['dim_indices'] == len(tight['pos_dims']) + len(tight['neg_dims']) def test_leaderboard_prefixes_match_the_shipped_lists(classifier): board = load('discovery/variant_leaderboard.json') assert board['top_pos_dims_30'][:N_POS] == classifier['pos_dims'] assert board['top_neg_dims_30'][:N_NEG] == classifier['neg_dims'] def test_tight_fpr_extra_dims_come_from_the_prop_sweep(tight): sweep = load('discovery/prop_specificity.json') k = len(tight['neg_dims_extra']) row = next(s for s in sweep['sweeps'] if s['extra_neg_k'] == k) assert row['added_dims'] == tight['neg_dims_extra'] assert abs(row['threshold'] - tight['threshold']) < 1e-6 def test_per_dim_thresholds_index_the_same_dims(classifier): cal = load('per_dim_thresholds.json') entries = cal['per_dim_thresholds'] assert [e['dim_global'] for e in entries] == classifier['pos_dims'] + classifier['neg_dims'] assert [e['dim_index_in_40'] for e in entries] == list(range(N_POS + N_NEG)) assert [e['is_pos'] for e in entries] == [True] * N_POS + [False] * N_NEG def test_quantized_thresholds_match_their_floats(): cal = load('per_dim_thresholds.json') scale = cal['quant_scale'] for e in cal['per_dim_thresholds']: assert e['threshold_int8'] == round(e['threshold'] * scale) assert -128 <= e['threshold_int8'] <= 127 def test_safetensors_agree_with_the_json_configs(classifier, tight): for name, c in (('classifier', classifier), ('classifier_tight_fpr', tight)): t = load_file(str(REPO / f'{name}.safetensors')) assert t['pos_dims'].tolist() == c['pos_dims'] assert t['neg_dims'].tolist() == c['neg_dims'] assert t['retained_dims'].tolist() == c['pos_dims'] + c['neg_dims'] assert abs(float(t['threshold'][0]) - c['threshold']) < 1e-4 w = t['retained_weight'][0].tolist() assert w == [1.0] * len(c['pos_dims']) + [-1.0] * len(c['neg_dims'])