"""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)