"""Unit tests for the scale-aware hybrid-LSE ABMIL engine.""" from __future__ import annotations from pathlib import Path import numpy as np import pytest import torch from src.core.services.abmil_model import ABMIL, ABMILInferenceEngine INPUT_DIM = 384 SCALE_EMBED_DIM = 16 NUM_SCALES = 3 def _write_checkpoint(tmp_path: Path, threshold: float = 0.3333) -> Path: """Build a fresh new-architecture ABMIL and persist it as a production-style checkpoint.""" model = ABMIL( input_dim=INPUT_DIM, attn_dim=256, classifier_hidden=256, attn_hidden=128, dropout=0.40, r=5.0, scale_embed_dim=SCALE_EMBED_DIM, num_scales=NUM_SCALES, ) model.eval() pt = tmp_path / "production_model.pt" torch.save( { "model_state_dict": model.state_dict(), "epoch": 75, "optimal_threshold": threshold, "input_dim": INPUT_DIM, "config": {"attn_dim": 256, "classifier_hidden": 256, "attn_hidden": 128, "dropout": 0.4}, }, pt, ) return pt def test_engine_loads_new_architecture(tmp_path: Path) -> None: pt = _write_checkpoint(tmp_path) engine = ABMILInferenceEngine(checkpoint_path=str(pt), device="cpu") engine.load() assert engine.loaded assert engine.metadata is not None assert engine.metadata.input_dim == INPUT_DIM assert engine.metadata.scale_embed_dim == SCALE_EMBED_DIM assert engine.metadata.num_scales == NUM_SCALES assert engine.metadata.lse_r == 5.0 # not stored in checkpoint -> falls back to settings default assert engine.metadata.threshold == pytest.approx(0.3333, abs=1e-4) def test_predict_embeddings_consumes_scale_indices(tmp_path: Path) -> None: pt = _write_checkpoint(tmp_path) engine = ABMILInferenceEngine(checkpoint_path=str(pt), device="cpu") engine.load() rng = np.random.default_rng(0) n_tiles = 12 embeddings = rng.standard_normal((n_tiles, INPUT_DIM)).astype(np.float32) scale_indices = np.array([0, 1] * (n_tiles // 2), dtype=np.int64) out = engine.predict_embeddings(embeddings, scale_indices) assert out["prediction"] in {"Pure", "Adulterated"} assert 0.0 <= out["probability_adulterated"] <= 1.0 assert out["num_tiles"] == n_tiles assert out["attention"].shape == (n_tiles,) assert out["threshold"] == pytest.approx(0.3333, abs=1e-4) def test_predict_requires_scale_indices_for_scale_aware_model(tmp_path: Path) -> None: pt = _write_checkpoint(tmp_path) engine = ABMILInferenceEngine(checkpoint_path=str(pt), device="cpu") engine.load() embeddings = np.zeros((4, INPUT_DIM), dtype=np.float32) with pytest.raises(ValueError, match="scale_indices are required"): engine.predict_embeddings(embeddings, None) def test_predict_rejects_wrong_embedding_dim(tmp_path: Path) -> None: pt = _write_checkpoint(tmp_path) engine = ABMILInferenceEngine(checkpoint_path=str(pt), device="cpu") engine.load() bad = np.zeros((4, 256), dtype=np.float32) with pytest.raises(RuntimeError, match="Embedding dimension mismatch"): engine.predict_embeddings(bad, np.zeros((4,), dtype=np.int64)) def test_predict_is_permutation_invariant(tmp_path: Path) -> None: """Attention + LSE pooling are set operations: tile order must not change the prediction.""" pt = _write_checkpoint(tmp_path) engine = ABMILInferenceEngine(checkpoint_path=str(pt), device="cpu") engine.load() rng = np.random.default_rng(7) embeddings = rng.standard_normal((10, INPUT_DIM)).astype(np.float32) scales = np.array([0, 1, 0, 1, 0, 1, 0, 1, 0, 1], dtype=np.int64) base = engine.predict_embeddings(embeddings, scales) perm = rng.permutation(10) shuffled = engine.predict_embeddings(embeddings[perm], scales[perm]) assert base["probability_adulterated"] == pytest.approx( shuffled["probability_adulterated"], abs=1e-5 )