cleanlab--cleanlab
a95a447010
* issuemanager.get_summary->make_summary * test(datalab): ♻️ move tests for issue managers into their respective test modules * test(datalab): ✅ add tests for custom issue manager - Validate scores provided to IssueManager.make_summary - Move fixture for custom issue manager to conftest.py - Test make_summary on a "custom" IssueManager * test(datalab): add __init__.py for discoverability of issue manager tests --------- Co-authored-by: Elías Snorrason <eliassno@gmail.com>
83 行
2.9 KiB
Python
83 行
2.9 KiB
Python
import numpy as np
|
|
import pytest
|
|
|
|
from cleanlab.datalab.issue_manager.duplicate import NearDuplicateIssueManager
|
|
|
|
SEED = 42
|
|
|
|
|
|
class TestNearDuplicateIssueManager:
|
|
@pytest.fixture
|
|
def embeddings(self, lab):
|
|
np.random.seed(SEED)
|
|
embeddings_array = 0.5 + 0.1 * np.random.rand(lab.get_info("statistics")["num_examples"], 2)
|
|
embeddings_array[4, :] = (
|
|
embeddings_array[3, :] + np.random.rand(embeddings_array.shape[1]) * 0.001
|
|
)
|
|
return {"embedding": embeddings_array}
|
|
|
|
@pytest.fixture
|
|
def issue_manager(self, lab, embeddings, monkeypatch):
|
|
mock_data = lab.data.from_dict({**lab.data.to_dict(), **embeddings})
|
|
monkeypatch.setattr(lab, "data", mock_data)
|
|
return NearDuplicateIssueManager(
|
|
datalab=lab,
|
|
metric="euclidean",
|
|
k=2,
|
|
)
|
|
|
|
def test_init(self, lab, issue_manager):
|
|
assert issue_manager.datalab == lab
|
|
assert issue_manager.metric == "euclidean"
|
|
assert issue_manager.k == 2
|
|
assert issue_manager.threshold == 0.13
|
|
|
|
issue_manager = NearDuplicateIssueManager(
|
|
datalab=lab,
|
|
threshold=0.1,
|
|
)
|
|
assert issue_manager.threshold == 0.1
|
|
|
|
def test_find_issues(self, issue_manager, embeddings):
|
|
issue_manager.find_issues(features=embeddings["embedding"])
|
|
issues, summary, info = issue_manager.issues, issue_manager.summary, issue_manager.info
|
|
expected_issue_mask = np.array([False] * 3 + [True] * 2)
|
|
assert np.all(
|
|
issues["is_near_duplicate_issue"] == expected_issue_mask
|
|
), "Issue mask should be correct"
|
|
assert summary["issue_type"][0] == "near_duplicate"
|
|
assert summary["score"][0] == pytest.approx(expected=0.03122489, abs=1e-7)
|
|
|
|
assert (
|
|
info.get("near_duplicate_sets", None) is not None
|
|
), "Should have sets of near duplicates"
|
|
|
|
new_issue_manager = NearDuplicateIssueManager(
|
|
datalab=issue_manager.datalab,
|
|
metric="euclidean",
|
|
k=2,
|
|
threshold=0.1,
|
|
)
|
|
new_issue_manager.find_issues(features=embeddings["embedding"])
|
|
|
|
def test_report(self, issue_manager, embeddings):
|
|
issue_manager.find_issues(features=embeddings["embedding"])
|
|
report = issue_manager.report(
|
|
issues=issue_manager.issues,
|
|
summary=issue_manager.summary,
|
|
info=issue_manager.info,
|
|
)
|
|
assert isinstance(report, str)
|
|
assert (
|
|
"------------------ near_duplicate issues -------------------\n\n"
|
|
"Number of examples with this issue:"
|
|
) in report
|
|
|
|
report = issue_manager.report(
|
|
issues=issue_manager.issues,
|
|
summary=issue_manager.summary,
|
|
info=issue_manager.info,
|
|
verbosity=3,
|
|
)
|
|
assert "Additional Information: " in report
|