""" Тесты предобработки, метрик и обучения. Метрики и порог проверяются на синтетических данных: важно, чтобы правило решения работало при сильном дисбалансе классов и чтобы подбор порога не деградировал, когда модель разделяет выборку почти идеально. """ import sys from pathlib import Path import numpy as np import pytest import torch sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from src.dxa.model import compute_metrics, create_model, select_threshold # noqa: E402 from src.dxa.preprocess import ( # noqa: E402 PreprocessConfig, normalize_array, preprocess_from_array, to_model_tensor, ) DATASET_ROOT = Path("dataset_hack") HAVE_DATASET = (DATASET_ROOT / "НД_для_обучения" / "Исследования").is_dir() needs_dataset = pytest.mark.skipif(not HAVE_DATASET, reason="dataset_hack is not available") class TestPreprocessConfig: def test_defaults_are_valid(self): cfg = PreprocessConfig() assert cfg.norm in ("percentile", "minmax") assert cfg.imagenet_norm is True def test_rejects_unknown_norm(self): with pytest.raises(ValueError): PreprocessConfig(norm="zscore") def test_roundtrip_through_dict(self): cfg = PreprocessConfig(norm="minmax", imagenet_norm=False, input_size=128) assert PreprocessConfig.from_dict(cfg.to_dict()) == cfg def test_from_dict_ignores_unknown_keys(self): cfg = PreprocessConfig.from_dict({"norm": "minmax", "legacy_field": 1}) assert cfg.norm == "minmax" def test_legacy_config_disables_imagenet_norm(self): assert PreprocessConfig.legacy().imagenet_norm is False class TestNormalizeArray: def test_minmax_maps_to_unit_range(self): arr = np.linspace(-100, 500, 1000).reshape(50, 20) out = normalize_array(arr, PreprocessConfig(norm="minmax")) assert out.min() == pytest.approx(0.0) assert out.max() == pytest.approx(1.0) assert out.dtype == np.float32 def test_percentile_clips_outliers(self): arr = np.zeros((100, 100), dtype=np.float32) arr[0, 0] = 10_000.0 # выброс out = normalize_array(arr, PreprocessConfig(norm="percentile", p_low=1, p_high=99)) assert out.max() <= 1.0 assert out.min() >= 0.0 def test_constant_image_does_not_divide_by_zero(self): out = normalize_array(np.full((10, 10), 7.0, dtype=np.float32), PreprocessConfig()) assert np.isfinite(out).all() class TestToModelTensor: def test_shape_and_channels(self): cfg = PreprocessConfig(input_size=64) out = to_model_tensor(np.zeros((30, 20), dtype=np.float32), cfg) assert out.shape == (3, 64, 64) assert out.dtype == np.float32 def test_imagenet_norm_changes_scale(self): # Тёмный пиксель: при стандартизации по ImageNet он становится заметно # отрицательным, поскольку среднее ImageNet для каналов ≈0.45. arr = np.full((16, 16), 0.2, dtype=np.float32) plain = to_model_tensor(arr, PreprocessConfig(input_size=16, imagenet_norm=False)) normed = to_model_tensor(arr, PreprocessConfig(input_size=16, imagenet_norm=True)) assert not np.allclose(plain, normed) assert normed.mean() < 0.0 # Без стандартизации значение остаётся тёмным, но положительным. assert 0.0 < plain.mean() < 0.25 def test_preprocess_from_array_matches_shape_contract(self): out = preprocess_from_array(np.random.rand(40, 30).astype(np.float32), PreprocessConfig(input_size=32)) assert out.shape == (3, 32, 32) class TestComputeMetrics: def test_perfect_separation(self): logits = torch.tensor([-5.0, -4.0, 4.0, 5.0]) labels = torch.tensor([0, 0, 1, 1]) m = compute_metrics(logits, labels, threshold=0.0) assert m["f1"] == pytest.approx(1.0) assert m["roc_auc"] == pytest.approx(1.0) assert m["tp"] == 2 and m["tn"] == 2 and m["fp"] == 0 and m["fn"] == 0 def test_inverted_predictions_give_zero_recall(self): logits = torch.tensor([5.0, 4.0, -4.0, -5.0]) labels = torch.tensor([0, 0, 1, 1]) m = compute_metrics(logits, labels, threshold=0.0) assert m["recall"] == 0.0 assert m["tp"] == 0 assert m["roc_auc"] == pytest.approx(0.0) def test_single_class_labels_yield_none_auc(self): m = compute_metrics(torch.tensor([-1.0, 1.0]), torch.tensor([0, 0])) assert m["roc_auc"] is None and m["pr_auc"] is None def test_threshold_is_reported_as_probability(self): m = compute_metrics(torch.tensor([0.0]), torch.tensor([0]), threshold=0.0) assert m["threshold_logit"] == 0.0 assert m["threshold_prob"] == pytest.approx(0.5) def test_counts_sum_to_sample_size(self): rng = np.random.default_rng(0) logits = torch.tensor(rng.normal(size=200)) labels = torch.tensor((rng.random(200) > 0.7).astype(int)) m = compute_metrics(logits, labels, threshold=0.0) assert m["tp"] + m["tn"] + m["fp"] + m["fn"] == 200 class TestSelectThreshold: def test_finds_separating_threshold(self): logits = torch.tensor([-3.0, -2.0, -1.0, 1.0, 2.0, 3.0]) labels = torch.tensor([0, 0, 0, 1, 1, 1]) threshold, m = select_threshold(logits, labels) assert m["f1"] == pytest.approx(1.0) assert -1.0 < threshold <= 1.0 def test_handles_imbalanced_data(self): """При ~10 % позитивов порог 0 даёт нулевой recall; подбор должен это исправить.""" rng = np.random.default_rng(1) logits = torch.tensor(np.concatenate([rng.normal(-1, 1, 90), rng.normal(0.5, 1, 10)])) labels = torch.tensor([0] * 90 + [1] * 10) _, m = select_threshold(logits, labels) assert m["recall"] > 0.5 def test_min_recall_constraint_is_respected(self): rng = np.random.default_rng(2) logits = torch.tensor(rng.normal(size=200)) labels = torch.tensor((rng.random(200) > 0.8).astype(int)) _, m = select_threshold(logits, labels, min_recall=0.9) assert m["recall"] >= 0.9 def test_single_class_returns_default(self): threshold, m = select_threshold(torch.tensor([1.0, 2.0]), torch.tensor([1, 1])) assert threshold == 0.0 assert m["roc_auc"] is None def test_unreachable_min_recall_falls_back(self): """Если требуемый recall недостижим, подбор не должен падать.""" logits = torch.tensor([5.0, 6.0, 7.0, 8.0]) labels = torch.tensor([0, 0, 1, 1]) _, m = select_threshold(logits, labels, min_recall=0.99) assert m["threshold_logit"] is not None class TestModelContract: def test_forward_returns_expected_keys_and_shapes(self): model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear") out = model.model(torch.zeros(2, 3, 224, 224)) assert set(out) == {"quality_logits", "region_logits"} assert out["quality_logits"].shape == (2, 2) assert out["region_logits"].shape == (2, 4) def test_linear_head_has_single_layer(self): model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear") linears = [m for m in model.model.quality_head if isinstance(m, torch.nn.Linear)] assert len(linears) == 1 def test_save_load_roundtrip_preserves_predictions(self, tmp_path): cfg = PreprocessConfig(input_size=224) model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear", learning_rate=1e-3) x = torch.randn(2, 3, 224, 224) with torch.no_grad(): before = model.model(x)["quality_logits"].clone() path = tmp_path / "ckpt.pth" model.save(path, preprocess=cfg, threshold=-1.5, epoch=3) reloaded = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear") metadata = reloaded.load(path) with torch.no_grad(): after = reloaded.model(x)["quality_logits"] assert torch.allclose(before, after, atol=1e-5) assert metadata["threshold"] == -1.5 assert metadata["epoch"] == 3 assert metadata["backbone"] == "resnet18" assert metadata["head"] == "linear" assert metadata["preprocess"] == cfg.to_dict() def test_feature_norm_buffers_survive_roundtrip(self, tmp_path): """Стандартизация признаков должна восстанавливаться вместе с весами.""" model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear") model.fit_feature_norm(_constant_loader(3)) x = torch.randn(2, 3, 224, 224) model.model.eval() with torch.no_grad(): before = model.model(x)["quality_logits"].clone() path = tmp_path / "ckpt.pth" model.save(path) reloaded = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear") reloaded.load(path) reloaded.model.eval() assert torch.allclose(reloaded.model.feat_mean, model.model.feat_mean) assert torch.allclose(reloaded.model.feat_std, model.model.feat_std) with torch.no_grad(): after = reloaded.model(x)["quality_logits"] assert torch.allclose(before, after, atol=1e-5) def test_frozen_backbone_keeps_batchnorm_in_eval(self): """ При заморозке backbone слои BatchNorm должны остаться в eval, иначе бегущие статистики сдвигаются и признаки расходятся с теми, на которых оценивалась стандартизация. """ model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear") model.set_backbone_trainable(False) model.train_epoch(_constant_loader(4)) bn_modules = [m for m in model.model.features.modules() if isinstance(m, torch.nn.modules.batchnorm._BatchNorm)] assert bn_modules, "resnet18 must contain BatchNorm layers" assert all(not m.training for m in bn_modules) def test_unfrozen_backbone_enables_batchnorm_training(self): model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear") model.set_backbone_trainable(True) model.train_epoch(_constant_loader(4)) bn_modules = [m for m in model.model.features.modules() if isinstance(m, torch.nn.modules.batchnorm._BatchNorm)] assert all(m.training for m in bn_modules) def test_mismatched_backbone_raises_on_load(self, tmp_path): model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear") path = tmp_path / "ckpt.pth" model.save(path) other = create_model(backbone="resnet34", pretrained=False, device="cpu", head="linear") with pytest.raises(RuntimeError): other.load(path) def test_raw_state_dict_is_rejected(self, tmp_path): model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear") path = tmp_path / "raw.pth" torch.save(model.model.state_dict(), path) with pytest.raises(ValueError): model.load(path) def test_fit_feature_norm_produces_finite_std(self): model = create_model(backbone="resnet18", pretrained=False, device="cpu", head="linear") model.fit_feature_norm(_constant_loader(4)) std = model.model.feat_std assert torch.isfinite(std).all() assert (std > 0).all(), "std must be strictly positive to avoid division by zero" class _SingleBatchLoader: """Минимальный лоадер из одного батча: нужен для проверок без датасета.""" def __init__(self, n): self.images = torch.zeros(n, 3, 224, 224) self.meta = { "label": torch.zeros(n, dtype=torch.long), "region_id": torch.zeros(n, dtype=torch.long), } self.dataset = range(n) def __iter__(self): yield self.images, self.meta def __len__(self): return 1 def _constant_loader(n): """Лоадер с постоянным входом: удобен для проверки статистик признаков.""" return _SingleBatchLoader(n) @needs_dataset class TestOnRealImages: def test_preprocess_dicom_is_deterministic(self): from src.dxa.labels import scan_dataset from src.dxa.preprocess import preprocess_dicom records = scan_dataset(DATASET_ROOT, with_pixel_dedup=True) cfg = PreprocessConfig(input_size=224) first = preprocess_dicom(records[0].path, cfg) second = preprocess_dicom(records[0].path, cfg) assert np.array_equal(first, second) assert first.shape == (3, 224, 224) assert np.isfinite(first).all() def test_duplicate_files_produce_identical_tensors(self): """Файлы с одинаковым содержимым должны давать одинаковый вход модели.""" from src.dxa.labels import scan_dataset from src.dxa.preprocess import preprocess_dicom records = scan_dataset(DATASET_ROOT, with_pixel_dedup=True) cfg = PreprocessConfig(input_size=224) duplicated = [r for r in records if len(r.sources) > 1] assert duplicated, "dataset should contain duplicates to test against" rec = duplicated[0] tensors = [preprocess_dicom(p, cfg) for p in rec.sources] for other in tensors[1:]: assert np.array_equal(tensors[0], other)