项目文件夹

文件
Jonas Mueller a95a447010 Update tests for custom issue manager example (#692)
* 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>
2023-05-03 00:40:27 +00:00

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