From 7e698360ff233731b25f50c25787ea3572f9b472 Mon Sep 17 00:00:00 2001 From: Trenton H <797416+stumpylog@users.noreply.github.com> Date: Tue, 15 Sep 2026 20:40:49 -0700 Subject: [PATCH] Performance: Drops fields from the MLPClassifier that are training only before pickling (#14114) --- src/documents/classifier.py | 22 +++++++++++++++++++++- 1 file changed, 21 insertions(+), 1 deletion(-) diff --git a/src/documents/classifier.py b/src/documents/classifier.py index 1e078c9c9..92d141489 100644 --- a/src/documents/classifier.py +++ b/src/documents/classifier.py @@ -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(