| 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) |
|
|