320 lines
14 KiB
Python
320 lines
14 KiB
Python
"""
|
||
Тесты предобработки, метрик и обучения.
|
||
|
||
Метрики и порог проверяются на синтетических данных: важно, чтобы правило
|
||
решения работало при сильном дисбалансе классов и чтобы подбор порога не
|
||
деградировал, когда модель разделяет выборку почти идеально.
|
||
"""
|
||
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)
|