CharlesCNorton
Image-level person classification on EUPE-ViT-B features with a single free parameter
f5498f9 | """The fused linear head is identical to the ternary sum it replaces.""" | |
| import sys | |
| import torch | |
| from common import D, score | |
| from conftest import REPO | |
| sys.path.insert(0, str(REPO)) | |
| from head import FusedClassifier # noqa: E402 | |
| def test_fused_head_equals_ternary_score(classifier): | |
| torch.manual_seed(0) | |
| model = FusedClassifier.from_config(None, REPO / 'classifier.json').eval() | |
| pooled = torch.randn(64, D) * 4.0 | |
| fused, present = model.head(pooled) | |
| direct = score(pooled, classifier['pos_dims'], classifier['neg_dims']) | |
| assert torch.allclose(fused, direct, atol=1e-4) | |
| assert torch.equal(present, direct > classifier['threshold']) | |
| def test_only_the_threshold_is_learnable(classifier): | |
| model = FusedClassifier.from_config(None, REPO / 'classifier.json') | |
| learnable = [n for n, p in model.named_parameters() if p.requires_grad] | |
| assert learnable == ['threshold'] | |
| assert sum(p.numel() for p in model.parameters()) == 1 | |
| def test_tight_fpr_head_reads_its_own_dim_count(tight): | |
| model = FusedClassifier.from_config( | |
| None, REPO / 'classifier_tight_fpr.json') | |
| assert model.retained_dims.numel() == len(tight['pos_dims']) + len(tight['neg_dims']) | |
| torch.manual_seed(1) | |
| pooled = torch.randn(32, D) * 4.0 | |
| fused, _ = model.head(pooled) | |
| direct = score(pooled, tight['pos_dims'], tight['neg_dims']) | |
| assert torch.allclose(fused, direct, atol=1e-4) | |