mirror of
https://github.com/paperless-ngx/paperless-ngx.git
synced 2026-09-12 04:37:58 +00:00
Enhancement: Improve matching for correspondents, storage path and labels by removing bias + adding minimum match threshold (#12164)
This commit is contained in:
+66
-28
@@ -34,6 +34,27 @@ from paperless.signed_pickle import signed_pickle_loads
|
||||
|
||||
logger = logging.getLogger("paperless.classifier")
|
||||
|
||||
|
||||
def _predict_with_threshold(classifier, X, threshold: float) -> int | None:
|
||||
"""
|
||||
Return the predicted class id, or None if:
|
||||
- the prediction is -1 (no match), or
|
||||
- the winning class probability is below the configured threshold.
|
||||
|
||||
Using predict_proba() instead of predict() lets us apply a minimum-confidence
|
||||
cutoff so that uncertain predictions are discarded rather than assigned.
|
||||
"""
|
||||
probas = classifier.predict_proba(X)[0]
|
||||
best_idx = int(probas.argmax())
|
||||
best_class = int(classifier.classes_[best_idx])
|
||||
|
||||
if best_class == -1:
|
||||
return None
|
||||
if threshold > 0.0 and probas[best_idx] < threshold:
|
||||
return None
|
||||
return best_class
|
||||
|
||||
|
||||
ADVANCED_TEXT_PROCESSING_ENABLED = (
|
||||
settings.NLTK_LANGUAGE is not None and settings.NLTK_ENABLED
|
||||
)
|
||||
@@ -102,7 +123,8 @@ class DocumentClassifier:
|
||||
# v8 - Added storage path classifier
|
||||
# v9 - Changed from hashing to time/ids for re-train check
|
||||
# v10 - HMAC-signed model file
|
||||
FORMAT_VERSION = 10
|
||||
# v11 - Use sample_weight for balanced training; predict_proba with threshold
|
||||
FORMAT_VERSION = 11
|
||||
|
||||
HMAC_SIZE = 32 # SHA-256 digest length
|
||||
|
||||
@@ -324,6 +346,13 @@ class DocumentClassifier:
|
||||
from sklearn.preprocessing import LabelBinarizer
|
||||
from sklearn.preprocessing import MultiLabelBinarizer
|
||||
|
||||
# MLPClassifier does not support class_weight directly
|
||||
# (https://github.com/scikit-learn/scikit-learn/issues/9113), so we use
|
||||
# compute_sample_weight to balance classes during training and prevent
|
||||
# over-represented correspondents from dominating predictions.
|
||||
# https://scikit-learn.org/stable/modules/generated/sklearn.utils.class_weight.compute_sample_weight.html
|
||||
from sklearn.utils.class_weight import compute_sample_weight
|
||||
|
||||
# Step 2: vectorize data
|
||||
logger.debug("Vectorizing data...")
|
||||
notify("Vectorizing document content...")
|
||||
@@ -369,7 +398,7 @@ class DocumentClassifier:
|
||||
self.tags_binarizer = MultiLabelBinarizer()
|
||||
labels_tags_vectorized = self.tags_binarizer.fit_transform(labels_tags)
|
||||
|
||||
self.tags_classifier = MLPClassifier(tol=0.01)
|
||||
self.tags_classifier = MLPClassifier(tol=0.01, random_state=0)
|
||||
self.tags_classifier.fit(data_vectorized, labels_tags_vectorized)
|
||||
else:
|
||||
self.tags_classifier = None
|
||||
@@ -380,8 +409,12 @@ class DocumentClassifier:
|
||||
notify(
|
||||
f"Training correspondent classifier ({num_correspondents} correspondent(s))...",
|
||||
)
|
||||
self.correspondent_classifier = MLPClassifier(tol=0.01)
|
||||
self.correspondent_classifier.fit(data_vectorized, labels_correspondent)
|
||||
self.correspondent_classifier = MLPClassifier(tol=0.01, random_state=0)
|
||||
self.correspondent_classifier.fit(
|
||||
data_vectorized,
|
||||
labels_correspondent,
|
||||
sample_weight=compute_sample_weight("balanced", labels_correspondent),
|
||||
)
|
||||
else:
|
||||
self.correspondent_classifier = None
|
||||
logger.debug(
|
||||
@@ -393,8 +426,12 @@ class DocumentClassifier:
|
||||
notify(
|
||||
f"Training document type classifier ({num_document_types} type(s))...",
|
||||
)
|
||||
self.document_type_classifier = MLPClassifier(tol=0.01)
|
||||
self.document_type_classifier.fit(data_vectorized, labels_document_type)
|
||||
self.document_type_classifier = MLPClassifier(tol=0.01, random_state=0)
|
||||
self.document_type_classifier.fit(
|
||||
data_vectorized,
|
||||
labels_document_type,
|
||||
sample_weight=compute_sample_weight("balanced", labels_document_type),
|
||||
)
|
||||
else:
|
||||
self.document_type_classifier = None
|
||||
logger.debug(
|
||||
@@ -406,10 +443,11 @@ class DocumentClassifier:
|
||||
"Training storage paths classifier...",
|
||||
)
|
||||
notify(f"Training storage path classifier ({num_storage_paths} path(s))...")
|
||||
self.storage_path_classifier = MLPClassifier(tol=0.01)
|
||||
self.storage_path_classifier = MLPClassifier(tol=0.01, random_state=0)
|
||||
self.storage_path_classifier.fit(
|
||||
data_vectorized,
|
||||
labels_storage_path,
|
||||
sample_weight=compute_sample_weight("balanced", labels_storage_path),
|
||||
)
|
||||
else:
|
||||
self.storage_path_classifier = None
|
||||
@@ -546,24 +584,24 @@ class DocumentClassifier:
|
||||
def predict_correspondent(self, content: str) -> int | None:
|
||||
if self.correspondent_classifier:
|
||||
X = self._vectorize(content)
|
||||
correspondent_id = self.correspondent_classifier.predict(X)
|
||||
if correspondent_id != -1:
|
||||
return correspondent_id
|
||||
else:
|
||||
return None
|
||||
else:
|
||||
return None
|
||||
predicted_id = _predict_with_threshold(
|
||||
self.correspondent_classifier,
|
||||
X,
|
||||
settings.CLASSIFIER_MATCH_THRESHOLD,
|
||||
)
|
||||
return predicted_id
|
||||
return None
|
||||
|
||||
def predict_document_type(self, content: str) -> int | None:
|
||||
if self.document_type_classifier:
|
||||
X = self._vectorize(content)
|
||||
document_type_id = self.document_type_classifier.predict(X)
|
||||
if document_type_id != -1:
|
||||
return document_type_id
|
||||
else:
|
||||
return None
|
||||
else:
|
||||
return None
|
||||
predicted_id = _predict_with_threshold(
|
||||
self.document_type_classifier,
|
||||
X,
|
||||
settings.CLASSIFIER_MATCH_THRESHOLD,
|
||||
)
|
||||
return predicted_id
|
||||
return None
|
||||
|
||||
def predict_tags(self, content: str) -> list[int]:
|
||||
from sklearn.utils.multiclass import type_of_target
|
||||
@@ -589,10 +627,10 @@ class DocumentClassifier:
|
||||
def predict_storage_path(self, content: str) -> int | None:
|
||||
if self.storage_path_classifier:
|
||||
X = self._vectorize(content)
|
||||
storage_path_id = self.storage_path_classifier.predict(X)
|
||||
if storage_path_id != -1:
|
||||
return storage_path_id
|
||||
else:
|
||||
return None
|
||||
else:
|
||||
return None
|
||||
predicted_id = _predict_with_threshold(
|
||||
self.storage_path_classifier,
|
||||
X,
|
||||
settings.CLASSIFIER_MATCH_THRESHOLD,
|
||||
)
|
||||
return predicted_id
|
||||
return None
|
||||
|
||||
Reference in New Issue
Block a user