DetectiveSAMv2 / tests /test_regression.py
Gertlek's picture
Clean DetectiveSAMv2-only release
0b2df65 verified
Raw
History Blame Contribute Delete
7.18 kB
from __future__ import annotations
import math
from pathlib import Path
import pytest
import torch
from detectivesam_inference.checkpoint import load_inference_config, resolve_checkpoint_path
from detectivesam_inference.dataset import PairDataset, prepare_sample
from detectivesam_inference.metrics import compute_f1, compute_iou, summarize_results
from detectivesam_inference.models.adapters import (
SpatialCrossAttentionSharedAdapter,
StreamEvidenceBuilder,
TransformerEvidenceMaskAdapter,
)
from detectivesam_inference.runtime import DetectiveSAMRunner, get_repo_root
def assert_close(value: float | None, expected: float, *, abs_tol: float = 1e-3) -> None:
assert value is not None
assert math.isclose(value, expected, rel_tol=0.0, abs_tol=abs_tol)
@pytest.fixture(scope="module")
def repo_root() -> Path:
return get_repo_root()
@pytest.fixture(scope="module")
def v2_runner(repo_root: Path) -> DetectiveSAMRunner:
checkpoint_path = resolve_checkpoint_path("detective_sam_v2", repo_root)
if not checkpoint_path.exists():
pytest.skip(f"Missing optional checkpoint: {checkpoint_path}")
if not torch.cuda.is_available():
pytest.skip("DetectiveSAMv2 regression metrics are tested with CUDA autocast.")
return DetectiveSAMRunner(checkpoint_path="detective_sam_v2", device="cuda")
def predict_metrics(
runner: DetectiveSAMRunner,
*,
source_path: Path,
target_path: Path,
mask_path: Path,
) -> tuple[float, float]:
sample = prepare_sample(
source_path=source_path,
target_path=target_path,
mask_path=mask_path,
img_size=runner.config.img_size,
perturbation_type=runner.config.perturbation_type,
perturbation_intensity=runner.config.perturbation_intensity,
)
prediction = runner.predict_sample(sample, threshold=0.5)
true_mask = sample.mask.squeeze().numpy().astype("uint8")
return compute_iou(prediction.pred_mask, true_mask), compute_f1(prediction.pred_mask, true_mask)
def test_checkpoint_alias_resolution(repo_root: Path) -> None:
assert resolve_checkpoint_path("detective_sam_v2", repo_root) == repo_root / "checkpoints" / "detective_sam_v2.pth"
def test_v2_checkpoint_sidecar(repo_root: Path) -> None:
config = load_inference_config(repo_root / "checkpoints" / "detective_sam_v2.pth")
assert config.prompt_dim == 96
assert config.downscale == 8
assert config.max_streams == 3
assert config.perturbation_type == "gaussian_blur+jpeg_compression+gaussian_noise"
assert config.adapter_type == "spatial_cross_attention"
assert config.mask_adapter_type == "transformer"
def test_json_checkpoint_sidecar(tmp_path: Path) -> None:
checkpoint_path = tmp_path / "best_model.pth"
checkpoint_path.touch()
(tmp_path / "model_params.json").write_text(
"""
{
"model_config": {
"prompt_dim": 96,
"downscale": 8,
"dropout_rate": 0.1,
"adapter_type": "spatial_cross_attention",
"mask_adapter_type": "transformer"
},
"training_config": {"img_size": 512},
"data_config": {
"perturbation_type": "gaussian_blur+jpeg_compression+gaussian_noise",
"perturbation_intensity": 0.5
},
"sam_config": {
"sam_config_file": "sam2.1_hiera_b+.yaml",
"sam_checkpoint": "sam2configs/sam2.1_hiera_base_plus.pt"
}
}
""",
encoding="utf-8",
)
config = load_inference_config(checkpoint_path)
assert config.prompt_dim == 96
assert config.max_streams == 3
assert config.adapter_type == "spatial_cross_attention"
assert config.mask_adapter_type == "transformer"
def test_legacy_architecture_config_is_rejected(tmp_path: Path) -> None:
checkpoint_path = tmp_path / "legacy.pth"
checkpoint_path.touch()
(tmp_path / "legacy_params.json").write_text(
"""
{
"model_config": {
"adapter_type": "conv",
"mask_adapter_type": "coarse"
}
}
""",
encoding="utf-8",
)
with pytest.raises(ValueError, match="DetectiveSAMv2-only"):
load_inference_config(checkpoint_path)
def test_adapter_exports_are_v2_only() -> None:
assert SpatialCrossAttentionSharedAdapter.__module__ == "detectivesam_inference.models.adapters"
assert StreamEvidenceBuilder.__module__ == "detectivesam_inference.models.adapters"
assert TransformerEvidenceMaskAdapter.__module__ == "detectivesam_inference.models.adapters"
def test_v2_banana_demo_metrics(repo_root: Path, v2_runner: DetectiveSAMRunner) -> None:
demo_root = repo_root / "demo" / "cocoglide"
iou, f1 = predict_metrics(
v2_runner,
source_path=demo_root / "source" / "banana_28809.png",
target_path=demo_root / "target" / "banana_28809.png",
mask_path=demo_root / "mask" / "banana_28809.png",
)
assert_close(iou, 0.8619145271101633)
assert_close(f1, 0.9258368357519847)
def test_v2_flux_demo_metrics(repo_root: Path, v2_runner: DetectiveSAMRunner) -> None:
demo_root = repo_root / "demo" / "flux_test"
iou, f1 = predict_metrics(
v2_runner,
source_path=demo_root / "source" / "548.png",
target_path=demo_root / "target" / "548.png",
mask_path=demo_root / "mask" / "548.png",
)
assert_close(iou, 0.8710592)
assert_close(f1, 0.9310867341877799)
def test_v2_qwen_demo_metrics(repo_root: Path, v2_runner: DetectiveSAMRunner) -> None:
demo_root = repo_root / "demo" / "qwen_test"
iou, f1 = predict_metrics(
v2_runner,
source_path=demo_root / "source" / "166.png",
target_path=demo_root / "target" / "166.png",
mask_path=demo_root / "mask" / "166.png",
)
assert_close(iou, 0.8415621398060885)
assert_close(f1, 0.9139655096240225)
def test_v2_cocoglide_eval_summary(repo_root: Path, v2_runner: DetectiveSAMRunner) -> None:
dataset = PairDataset(
root_dir=repo_root / "demo" / "cocoglide",
img_size=v2_runner.config.img_size,
perturbation_type=v2_runner.config.perturbation_type,
perturbation_intensity=v2_runner.config.perturbation_intensity,
)
per_sample_results: list[dict[str, float | str | None]] = []
for sample in dataset:
prediction = v2_runner.predict_sample(sample, threshold=0.5)
true_mask = sample.mask.squeeze().numpy().astype("uint8")
per_sample_results.append(
{
"name": sample.name,
"iou": compute_iou(prediction.pred_mask, true_mask),
"f1": compute_f1(prediction.pred_mask, true_mask),
}
)
summary = summarize_results(per_sample_results)
assert summary["num_samples"] == 2
assert summary["num_samples_with_gt"] == 2
assert_close(summary["mean_iou"], 0.7776422818026876)
assert_close(summary["mean_f1"], 0.8723800367494952)
expected_by_name = {
"banana_28809": (0.8619145271101633, 0.9258368357519847),
"train_221213": (0.693370036495212, 0.8189232377470056),
}
for result in per_sample_results:
expected_iou, expected_f1 = expected_by_name[result["name"]]
assert_close(result["iou"], expected_iou)
assert_close(result["f1"], expected_f1)