mirror of
https://github.com/paperless-ngx/paperless-ngx.git
synced 2026-09-10 19:58:01 +00:00
Enhancement: Improve matching for correspondents, storage path and labels by removing bias + adding minimum match threshold (#12164)
This commit is contained in:
@@ -3,6 +3,7 @@ import warnings
|
||||
from pathlib import Path
|
||||
from unittest import mock
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
from django.conf import settings
|
||||
from django.test import TestCase
|
||||
@@ -11,6 +12,7 @@ from django.test import override_settings
|
||||
from documents.classifier import ClassifierModelCorruptError
|
||||
from documents.classifier import DocumentClassifier
|
||||
from documents.classifier import IncompatibleClassifierVersionError
|
||||
from documents.classifier import _predict_with_threshold
|
||||
from documents.classifier import load_classifier
|
||||
from documents.models import Correspondent
|
||||
from documents.models import Document
|
||||
@@ -625,6 +627,103 @@ class TestClassifier(DirectoriesMixin, TestCase):
|
||||
self.assertEqual(self.classifier.predict_storage_path(doc1.content), sp.pk)
|
||||
self.assertIsNone(self.classifier.predict_storage_path(doc2.content))
|
||||
|
||||
def test_predict_rejects_prediction_below_match_threshold(self) -> None:
|
||||
"""
|
||||
GIVEN:
|
||||
- Classifiers trained against test data with confident predictions
|
||||
WHEN:
|
||||
- CLASSIFIER_MATCH_THRESHOLD exceeds the model's confidence
|
||||
THEN:
|
||||
- Every predict_* method discards the match in favor of no match
|
||||
"""
|
||||
c1 = Correspondent.objects.create(
|
||||
name="c1",
|
||||
matching_algorithm=Correspondent.MATCH_AUTO,
|
||||
)
|
||||
dt1 = DocumentType.objects.create(
|
||||
name="dt1",
|
||||
matching_algorithm=DocumentType.MATCH_AUTO,
|
||||
)
|
||||
sp1 = StoragePath.objects.create(
|
||||
name="sp1",
|
||||
matching_algorithm=StoragePath.MATCH_AUTO,
|
||||
)
|
||||
|
||||
doc1 = Document.objects.create(
|
||||
title="doc1",
|
||||
content="this is a document from c1",
|
||||
correspondent=c1,
|
||||
document_type=dt1,
|
||||
storage_path=sp1,
|
||||
checksum="A",
|
||||
)
|
||||
Document.objects.create(
|
||||
title="doc2",
|
||||
content="this is a document from no one",
|
||||
checksum="B",
|
||||
)
|
||||
|
||||
self.classifier.train()
|
||||
|
||||
predictors = {
|
||||
"correspondent": self.classifier.predict_correspondent,
|
||||
"document_type": self.classifier.predict_document_type,
|
||||
"storage_path": self.classifier.predict_storage_path,
|
||||
}
|
||||
# No real prediction can reach a confidence this high, so this
|
||||
# isolates the threshold check from the model's actual output.
|
||||
with override_settings(CLASSIFIER_MATCH_THRESHOLD=0.999999):
|
||||
for name, predict in predictors.items():
|
||||
with self.subTest(field=name):
|
||||
self.assertIsNone(predict(doc1.content))
|
||||
|
||||
def test_train_uses_balanced_sample_weight(self) -> None:
|
||||
"""
|
||||
GIVEN:
|
||||
- A training set with correspondents, document types and storage paths
|
||||
WHEN:
|
||||
- The classifier is trained
|
||||
THEN:
|
||||
- Each MLP classifier is fit with balanced sample weights, so that
|
||||
over-represented classes don't dominate predictions
|
||||
"""
|
||||
c1 = Correspondent.objects.create(
|
||||
name="c1",
|
||||
matching_algorithm=Correspondent.MATCH_AUTO,
|
||||
)
|
||||
dt1 = DocumentType.objects.create(
|
||||
name="dt1",
|
||||
matching_algorithm=DocumentType.MATCH_AUTO,
|
||||
)
|
||||
sp1 = StoragePath.objects.create(
|
||||
name="sp1",
|
||||
matching_algorithm=StoragePath.MATCH_AUTO,
|
||||
)
|
||||
|
||||
Document.objects.create(
|
||||
title="doc1",
|
||||
content="this is a document from c1",
|
||||
correspondent=c1,
|
||||
document_type=dt1,
|
||||
storage_path=sp1,
|
||||
checksum="A",
|
||||
)
|
||||
Document.objects.create(
|
||||
title="doc2",
|
||||
content="this is a document from no one",
|
||||
checksum="B",
|
||||
)
|
||||
|
||||
with mock.patch(
|
||||
"sklearn.utils.class_weight.compute_sample_weight",
|
||||
return_value=None,
|
||||
) as mocked_compute_sample_weight:
|
||||
self.classifier.train()
|
||||
|
||||
self.assertEqual(mocked_compute_sample_weight.call_count, 3)
|
||||
for call in mocked_compute_sample_weight.call_args_list:
|
||||
self.assertEqual(call.args[0], "balanced")
|
||||
|
||||
def test_one_tag_predict(self) -> None:
|
||||
t1 = Tag.objects.create(name="t1", matching_algorithm=Tag.MATCH_AUTO, pk=12)
|
||||
|
||||
@@ -810,6 +909,52 @@ class TestClassifier(DirectoriesMixin, TestCase):
|
||||
load_classifier(raise_exception=True)
|
||||
|
||||
|
||||
class _StubProbaClassifier:
|
||||
"""
|
||||
A fake scikit-learn classifier exposing just enough of the API for
|
||||
`_predict_with_threshold`: `classes_` and `predict_proba`.
|
||||
"""
|
||||
|
||||
def __init__(self, classes: list[int], probabilities: list[float]) -> None:
|
||||
self.classes_ = np.array(classes)
|
||||
self._probabilities = np.array([probabilities])
|
||||
|
||||
def predict_proba(self, X) -> np.ndarray:
|
||||
return self._probabilities
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("classes", "probabilities", "threshold", "expected"),
|
||||
[
|
||||
# confident prediction above the threshold is returned
|
||||
([-1, 3], [0.1, 0.9], 0.6, 3),
|
||||
# prediction below the threshold is discarded
|
||||
([-1, 3], [0.45, 0.55], 0.6, None),
|
||||
# boundary: exactly at the threshold is accepted, not discarded
|
||||
([-1, 3], [0.4, 0.6], 0.6, 3),
|
||||
# the winning class is the "no match" pseudo-class, regardless of its
|
||||
# own confidence
|
||||
([-1, 3], [0.99, 0.01], 0.0, None),
|
||||
# threshold of 0.0 disables the confidence check entirely
|
||||
([-1, 3], [0.45, 0.55], 0.0, 3),
|
||||
],
|
||||
)
|
||||
def test_predict_with_threshold(classes, probabilities, threshold, expected) -> None:
|
||||
classifier = _StubProbaClassifier(classes, probabilities)
|
||||
result = _predict_with_threshold(classifier, X=None, threshold=threshold)
|
||||
assert result == expected
|
||||
|
||||
|
||||
def test_classifier_match_threshold_default() -> None:
|
||||
"""
|
||||
GIVEN:
|
||||
- No PAPERLESS_CLASSIFIER_MATCH_THRESHOLD environment variable is set
|
||||
THEN:
|
||||
- The classifier match threshold defaults to 0.6
|
||||
"""
|
||||
assert settings.CLASSIFIER_MATCH_THRESHOLD == 0.6
|
||||
|
||||
|
||||
def test_preprocess_content() -> None:
|
||||
"""
|
||||
GIVEN:
|
||||
|
||||
Reference in New Issue
Block a user