mirror of
https://github.com/paperless-ngx/paperless-ngx.git
synced 2026-09-16 06:38:00 +00:00
Performance: Drops fields from the MLPClassifier that are training only before pickling (#14114)
This commit is contained in:
@@ -18,6 +18,7 @@ if TYPE_CHECKING:
|
||||
from typing import Self
|
||||
|
||||
from numpy import ndarray
|
||||
from sklearn.neural_network import MLPClassifier
|
||||
|
||||
from django.conf import settings
|
||||
from django.core.cache import cache
|
||||
@@ -174,7 +175,8 @@ class DocumentClassifier:
|
||||
# v8 - Added storage path classifier
|
||||
# v9 - Changed from hashing to time/ids for re-train check
|
||||
# v10 - HMAC-signed model file
|
||||
# v11 - Use sample_weight for balanced training; predict_proba with threshold
|
||||
# v11 - Use sample_weight for balanced training; predict_proba with threshold;
|
||||
# drop training-only MLP state before saving
|
||||
FORMAT_VERSION = 11
|
||||
|
||||
HMAC_SIZE = 32 # SHA-256 digest length
|
||||
@@ -212,6 +214,20 @@ class DocumentClassifier:
|
||||
def _new_hmac() -> hmac.HMAC:
|
||||
return hmac.new(settings.SECRET_KEY.encode(), digestmod=sha256)
|
||||
|
||||
@staticmethod
|
||||
def _strip_training_state(classifier: MLPClassifier) -> None:
|
||||
"""
|
||||
Drop MLPClassifier state which is only used during fit(), never by predict().
|
||||
|
||||
The Adam optimizer keeps two moment arrays the size of the weights, and
|
||||
without early_stopping _best_coefs/_best_intercepts are just a copy of the
|
||||
initial random weights. Together that is 3x the size of the weights, which
|
||||
would otherwise be pickled and loaded along with the model.
|
||||
"""
|
||||
del classifier._optimizer
|
||||
del classifier._best_coefs
|
||||
del classifier._best_intercepts
|
||||
|
||||
@staticmethod
|
||||
def _compute_hmac(data: bytes | memoryview) -> bytes:
|
||||
mac = DocumentClassifier._new_hmac()
|
||||
@@ -477,6 +493,7 @@ class DocumentClassifier:
|
||||
|
||||
self.tags_classifier = MLPClassifier(tol=0.01, random_state=0)
|
||||
self.tags_classifier.fit(data_vectorized, labels_tags_vectorized)
|
||||
self._strip_training_state(self.tags_classifier)
|
||||
else:
|
||||
self.tags_classifier = None
|
||||
logger.debug("There are no tags. Not training tags classifier.")
|
||||
@@ -492,6 +509,7 @@ class DocumentClassifier:
|
||||
labels_correspondent,
|
||||
sample_weight=compute_sample_weight("balanced", labels_correspondent),
|
||||
)
|
||||
self._strip_training_state(self.correspondent_classifier)
|
||||
else:
|
||||
self.correspondent_classifier = None
|
||||
logger.debug(
|
||||
@@ -509,6 +527,7 @@ class DocumentClassifier:
|
||||
labels_document_type,
|
||||
sample_weight=compute_sample_weight("balanced", labels_document_type),
|
||||
)
|
||||
self._strip_training_state(self.document_type_classifier)
|
||||
else:
|
||||
self.document_type_classifier = None
|
||||
logger.debug(
|
||||
@@ -526,6 +545,7 @@ class DocumentClassifier:
|
||||
labels_storage_path,
|
||||
sample_weight=compute_sample_weight("balanced", labels_storage_path),
|
||||
)
|
||||
self._strip_training_state(self.storage_path_classifier)
|
||||
else:
|
||||
self.storage_path_classifier = None
|
||||
logger.debug(
|
||||
|
||||
Reference in New Issue
Block a user