Skip to content

Commit b6b7257

Browse files
committed
Add unit tests for crop_generator and ds_utils utility functions
1 parent 366cb65 commit b6b7257

3 files changed

Lines changed: 184 additions & 0 deletions

File tree

Lines changed: 89 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,89 @@
1+
"""
2+
Tests for crop_generator.py utility functions.
3+
"""
4+
5+
import pytest
6+
7+
from virtual_stain_flow.datasets.base_dataset import BaseImageDataset
8+
from virtual_stain_flow.datasets.ds_engine.crop_generator import (
9+
_compute_center_crop,
10+
generate_center_crops,
11+
)
12+
13+
14+
class TestComputeCenterCrop:
15+
"""Tests for _compute_center_crop function."""
16+
17+
def test_computes_correct_center_for_even_dimensions(self):
18+
"""Should compute correct top-left for centered crop."""
19+
x, y = _compute_center_crop(100, 100, 50)
20+
assert (x, y) == (25, 25)
21+
22+
def test_computes_correct_center_for_odd_remainder(self):
23+
"""Should floor division for odd remainder."""
24+
x, y = _compute_center_crop(10, 10, 5)
25+
assert (x, y) == (2, 2)
26+
27+
def test_computes_correct_center_for_rectangular_image(self):
28+
"""Should handle non-square images."""
29+
x, y = _compute_center_crop(100, 50, 20)
30+
assert (x, y) == (40, 15)
31+
32+
def test_raises_error_when_crop_exceeds_width(self):
33+
"""Should raise ValueError when crop_size > width."""
34+
with pytest.raises(ValueError, match="exceeds image dimensions"):
35+
_compute_center_crop(10, 100, 20)
36+
37+
def test_raises_error_when_crop_exceeds_height(self):
38+
"""Should raise ValueError when crop_size > height."""
39+
with pytest.raises(ValueError, match="exceeds image dimensions"):
40+
_compute_center_crop(100, 10, 20)
41+
42+
43+
class TestGenerateCenterCrops:
44+
"""Tests for generate_center_crops function."""
45+
46+
def test_generates_center_crops_for_all_samples(self, basic_dataset):
47+
"""Should generate one center crop per sample."""
48+
crop_specs = generate_center_crops(basic_dataset, crop_size=4)
49+
50+
# Should have 3 samples (from file_index fixture)
51+
assert len(crop_specs) == 3
52+
assert set(crop_specs.keys()) == {0, 1, 2}
53+
54+
def test_crop_specs_have_correct_format(self, basic_dataset):
55+
"""Should return crop specs in ((x, y), width, height) format."""
56+
crop_specs = generate_center_crops(basic_dataset, crop_size=4)
57+
58+
for idx, crops in crop_specs.items():
59+
assert len(crops) == 1 # One center crop per sample
60+
(x, y), w, h = crops[0]
61+
assert w == 4
62+
assert h == 4
63+
# For 10x10 images with crop_size=4: center is at (3, 3)
64+
assert (x, y) == (3, 3)
65+
66+
def test_raises_error_for_non_positive_crop_size(self, basic_dataset):
67+
"""Should raise ValueError for crop_size <= 0."""
68+
with pytest.raises(ValueError, match="crop_size must be positive"):
69+
generate_center_crops(basic_dataset, crop_size=0)
70+
71+
with pytest.raises(ValueError, match="crop_size must be positive"):
72+
generate_center_crops(basic_dataset, crop_size=-1)
73+
74+
def test_raises_error_when_no_active_channels(self, file_index):
75+
"""Should raise ValueError when no channels configured."""
76+
dataset = BaseImageDataset(
77+
file_index=file_index,
78+
pil_image_mode="I;16",
79+
input_channel_keys=None,
80+
target_channel_keys=None,
81+
)
82+
with pytest.raises(ValueError, match="No active channels"):
83+
generate_center_crops(dataset, crop_size=4)
84+
85+
def test_raises_error_when_crop_too_large(self, basic_dataset):
86+
"""Should raise ValueError when crop_size exceeds image dimensions."""
87+
# Images are 10x10, so crop_size=20 should fail
88+
with pytest.raises(ValueError, match="exceeds image dimensions"):
89+
generate_center_crops(basic_dataset, crop_size=20)
Lines changed: 70 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,70 @@
1+
"""
2+
Tests for ds_utils.py utility functions.
3+
"""
4+
5+
import pytest
6+
7+
from virtual_stain_flow.datasets.ds_engine.ds_utils import (
8+
_get_active_channels,
9+
_validate_same_dimensions_across_channels,
10+
)
11+
12+
13+
class TestGetActiveChannels:
14+
"""Tests for _get_active_channels function."""
15+
16+
def test_returns_union_of_input_and_target_channels(self, basic_dataset):
17+
"""Should return all unique channel keys from input and target."""
18+
result = _get_active_channels(basic_dataset)
19+
assert result == ["input_ch1", "input_ch2", "target_ch1"]
20+
21+
def test_removes_duplicates_preserving_order(self, file_index):
22+
"""Should remove duplicates while preserving order."""
23+
from virtual_stain_flow.datasets.base_dataset import BaseImageDataset
24+
25+
# Create dataset with overlapping channel keys
26+
dataset = BaseImageDataset(
27+
file_index=file_index,
28+
pil_image_mode="I;16",
29+
input_channel_keys=["input_ch1", "input_ch2"],
30+
target_channel_keys=["input_ch1"], # Duplicate
31+
)
32+
result = _get_active_channels(dataset)
33+
assert result == ["input_ch1", "input_ch2"]
34+
35+
36+
class TestValidateSameDimensionsAcrossChannels:
37+
"""Tests for _validate_same_dimensions_across_channels function."""
38+
39+
def test_returns_common_dimension_when_all_match(self):
40+
"""Should return the common dimension when all channels match."""
41+
dims = ((10, 10), (10, 10), (10, 10))
42+
channels = ["ch1", "ch2", "ch3"]
43+
result = _validate_same_dimensions_across_channels(dims, channels, idx=0)
44+
assert result == (10, 10)
45+
46+
def test_raises_error_on_dimension_mismatch(self):
47+
"""Should raise ValueError when dimensions don't match."""
48+
dims = ((10, 10), (20, 20))
49+
channels = ["ch1", "ch2"]
50+
with pytest.raises(ValueError, match="Dimension mismatch"):
51+
_validate_same_dimensions_across_channels(dims, channels, idx=0)
52+
53+
def test_raises_error_on_empty_dims(self):
54+
"""Should raise ValueError when dims is empty."""
55+
with pytest.raises(ValueError, match="No dimensions returned"):
56+
_validate_same_dimensions_across_channels((), [], idx=0)
57+
58+
def test_raises_error_when_all_files_missing(self):
59+
"""Should raise ValueError when all channel files are missing (None)."""
60+
dims = (None, None)
61+
channels = ["ch1", "ch2"]
62+
with pytest.raises(ValueError, match="All channel files missing"):
63+
_validate_same_dimensions_across_channels(dims, channels, idx=0)
64+
65+
def test_ignores_none_values_for_missing_files(self):
66+
"""Should skip None values and validate remaining dimensions."""
67+
dims = ((10, 10), None, (10, 10))
68+
channels = ["ch1", "ch2", "ch3"]
69+
result = _validate_same_dimensions_across_channels(dims, channels, idx=0)
70+
assert result == (10, 10)

‎tests/datasets/test_crop_dataset.py‎

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -336,3 +336,28 @@ def test_dataloader_with_shuffle(self, crop_dataset):
336336
for inp_batch, tar_batch in loader:
337337
assert inp_batch.shape[0] <= 2
338338
assert tar_batch.shape[0] <= 2
339+
340+
341+
class TestCropImageDatasetFromBaseDataset:
342+
"""Test suite for CropImageDataset.from_base_dataset class method."""
343+
344+
def test_from_base_dataset_default_center_crops(self, basic_dataset):
345+
"""Test from_base_dataset creates CropImageDataset with default center crops."""
346+
crop_ds = CropImageDataset.from_base_dataset(basic_dataset, crop_size=4)
347+
348+
# Should have one center crop per image (3 images)
349+
assert len(crop_ds) == 3
350+
351+
# Should preserve channel keys
352+
assert crop_ds.input_channel_keys == basic_dataset.input_channel_keys
353+
assert crop_ds.target_channel_keys == basic_dataset.target_channel_keys
354+
assert crop_ds.pil_image_mode == basic_dataset.pil_image_mode
355+
356+
# Verify crop shape and values
357+
inp_np, tar_np = crop_ds.get_raw_item(0)
358+
assert inp_np.shape == (2, 4, 4)
359+
assert tar_np.shape == (1, 4, 4)
360+
361+
# Center crop of 10x10 image with crop_size=4 starts at (3, 3)
362+
assert crop_ds.crop_info.x == 3
363+
assert crop_ds.crop_info.y == 3

0 commit comments

Comments
 (0)