bone_2026/tests/test_preprocess_and_model.py

320 lines
14 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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