signature-verification / tests /test_prepare_data.py
imaflower's picture
Upload folder using huggingface_hub
4a95d4d verified
Raw
History Blame Contribute Delete
3.37 kB
import os
from unittest.mock import mock_open, patch
import numpy as np
import pytest
from signature_verification.prepare_data import generate_batch, get_groups
IMG_H = 155
IMG_W = 220
PATH = "test_path"
@pytest.fixture
def read_file(monkeypatch):
# Fake content of data_dir_list which contain the index of the writers
data_dir_content = "37,46"
m_open = mock_open(read_data=data_dir_content)
monkeypatch.setattr("builtins.open", m_open)
@pytest.fixture
def fake_os_listdir(monkeypatch):
def fake_listdir(path):
dir1_files = [f"forgeries_37_{i}.png" for i in range(1, 25)] + [
f"original_37_{i}.png" for i in range(1, 25)
]
dir2_files = [f"forgeries_46_{i}.png" for i in range(1, 25)] + [
f"original_46_{i}.png" for i in range(1, 25)
]
if path.endswith("37"):
return dir1_files
elif path.endswith("46"):
return dir2_files
else:
return []
monkeypatch.setattr(os, "listdir", fake_listdir)
@pytest.fixture
def create_txt_file(tmp_path):
# Create a txt file to save idx_writers
data_dir_content = "37,46"
txt_path = tmp_path / "data_dir_list.txt"
txt_path.write_text(data_dir_content)
return txt_path
def test_create_file(tmp_path):
data_dir_content = "37,46"
p = tmp_path / "data_dir_list.txt"
p.write_text(data_dir_content)
assert p.read_text() == data_dir_content
assert len(list(tmp_path.iterdir())) == 1
def create_expected_groups():
expected_orig_groups = [
[f"{PATH}/37/original_37_{i}.png" for i in range(1, 25)],
[f"{PATH}/46/original_46_{i}.png" for i in range(1, 25)],
]
expected_forg_groups = [
[f"{PATH}/37/forgeries_37_{i}.png" for i in range(1, 25)],
[f"{PATH}/46/forgeries_46_{i}.png" for i in range(1, 25)],
]
return expected_orig_groups, expected_forg_groups
def test_get_groups(read_file, fake_os_listdir):
expected_orig_groups, expected_forg_groups = create_expected_groups()
for group in expected_orig_groups, expected_forg_groups:
group.sort()
sorted_expected_orig_groups = [sorted(group) for group in expected_orig_groups]
sorted_expected_forg_groups = [sorted(group) for group in expected_forg_groups]
actual_orig_groups, actual_forg_groups = get_groups(PATH, fake_os_listdir)
# Check if all original signatures in orig_groups and all forged signatures in forg_groups
assert sorted_expected_orig_groups == actual_orig_groups
assert sorted_expected_forg_groups == actual_forg_groups
def create_cv2_imread(filepath, flags):
return np.ones((IMG_H, IMG_W), dtype=np.uint8) * 255
def test_generate_batch(read_file, fake_os_listdir):
orig_groups, forg_groups = create_expected_groups()
with patch("cv2.imread", side_effect=create_cv2_imread):
batch_size = 128
actual_batch = generate_batch(orig_groups, forg_groups, batch_size)
pairs, targets = next(actual_batch)
# Check if pairs are generated by generate_batch have the shape as input
assert pairs[0].shape == (batch_size, IMG_H, IMG_W, 1)
assert pairs[1].shape == (batch_size, IMG_H, IMG_W, 1)
assert targets.shape == (batch_size,)