Coverage for documents/classifier.py: 16%
336 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 09:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 09:07 +0000
1from __future__ import annotations
3import functools
4import hmac
5import logging
6import pickle
7import re
8import unicodedata
9import warnings
10from hashlib import sha256
11from pathlib import Path
12from typing import TYPE_CHECKING
14if TYPE_CHECKING: 14 ↛ 15line 14 didn't jump to line 15 because the condition on line 14 was never true
15 from collections.abc import Callable
16 from collections.abc import Iterator
17 from datetime import datetime
18 from types import TracebackType
19 from typing import BinaryIO
20 from typing import Self
22 import tantivy
23 from numpy import ndarray
24 from sklearn.neural_network import MLPClassifier
26from django.conf import settings
27from django.core.cache import cache
28from django.core.cache import caches
29from django.db.models import Prefetch
31from documents._snowball_stopwords import ENGLISH as ENGLISH_STOP_WORDS
32from documents.caching import CACHE_5_MINUTES
33from documents.caching import CACHE_50_MINUTES
34from documents.caching import CLASSIFIER_HASH_KEY
35from documents.caching import CLASSIFIER_MODIFIED_KEY
36from documents.caching import CLASSIFIER_VERSION_KEY
37from documents.models import Document
38from documents.models import MatchingModel
39from documents.models import Tag
40from paperless.signed_pickle import SignedPickleError
41from paperless.signed_pickle import signed_pickle_dumps
42from paperless.signed_pickle import signed_pickle_loads
44logger = logging.getLogger("paperless.classifier")
47def _predict_with_threshold(classifier, X, threshold: float) -> int | None:
48 """
49 Return the predicted class id, or None if:
50 - the prediction is -1 (no match), or
51 - the winning class probability is below the configured threshold.
53 Using predict_proba() instead of predict() lets us apply a minimum-confidence
54 cutoff so that uncertain predictions are discarded rather than assigned.
55 """
56 probas = classifier.predict_proba(X)[0]
57 best_idx = int(probas.argmax())
58 best_class = int(classifier.classes_[best_idx])
60 if best_class == -1:
61 return None
62 if threshold > 0.0 and probas[best_idx] < threshold:
63 return None
64 return best_class
67read_cache = caches["read-cache"]
70RE_WORD = re.compile(r"\b[\w]+\b") # words that may contain digits
72# Documents whose content is fetched per query while training
73_CONTENT_CHUNK_SIZE = 1000
76class _SignedFileWriter:
77 """
78 Atomically writes a file made of an HMAC signature followed by the data,
79 signing the data as it streams to disk rather than holding it in memory.
81 The signature is only known once everything is written, so its space is
82 reserved at the start of the file and filled in on exit. The target is only
83 replaced once the file is complete; on error the partial file is removed.
84 """
86 def __init__(self, target: Path, mac: hmac.HMAC) -> None:
87 self._target = target
88 self._temp = target.with_name(f"{target.name}.part")
89 self._mac = mac
90 self._file: BinaryIO
92 def __enter__(self) -> Self:
93 self._file = self._temp.open("wb")
94 self._file.write(bytes(self._mac.digest_size))
95 return self
97 def write(self, data: bytes | memoryview) -> int:
98 self._mac.update(data)
99 return self._file.write(data)
101 def __exit__(
102 self,
103 exc_type: type[BaseException] | None,
104 exc_value: BaseException | None,
105 traceback: TracebackType | None,
106 ) -> None:
107 try:
108 with self._file:
109 if exc_type is None:
110 self._file.seek(0)
111 self._file.write(self._mac.digest())
112 if exc_type is None:
113 self._temp.rename(self._target)
114 finally:
115 # A no-op after a successful rename, otherwise removes the partial file
116 self._temp.unlink(missing_ok=True)
119@functools.cache
120def _text_analyzer(language: str) -> tantivy.TextAnalyzer:
121 """
122 Builds the cached analyzer for a language: word tokens, lowercase, stop words, stemmer.
124 Long tokens are kept and accents are not folded to ASCII, since stemmers
125 for languages such as French and German rely on accents.
126 """
127 import tantivy
129 if language == "english":
130 # Tantivy's builtin English list is much shorter than Snowball's.
131 # Split contractions ("don't") on word characters, as content is, so they match
132 stop_words = tantivy.Filter.custom_stopword(
133 sorted({t for word in ENGLISH_STOP_WORDS for t in RE_WORD.findall(word)}),
134 )
135 else:
136 stop_words = tantivy.Filter.stopword(language)
137 return (
138 tantivy.TextAnalyzerBuilder(tantivy.Tokenizer.regex(r"\w+"))
139 .filter(tantivy.Filter.lowercase())
140 .filter(stop_words)
141 .filter(tantivy.Filter.stemmer(language))
142 .build()
143 )
146class IncompatibleClassifierVersionError(Exception):
147 def __init__(self, message: str, *args: object) -> None:
148 self.message: str = message
149 super().__init__(*args)
152class ClassifierModelCorruptError(Exception):
153 pass
156def load_classifier(*, raise_exception: bool = False) -> DocumentClassifier | None:
157 if not settings.MODEL_FILE.is_file():
158 logger.debug(
159 "Document classification model does not exist (yet), not "
160 "performing automatic matching.",
161 )
162 return None
164 classifier = DocumentClassifier()
165 try:
166 classifier.load()
168 except IncompatibleClassifierVersionError as e:
169 logger.info(f"Classifier version incompatible: {e.message}, will re-train")
170 Path(settings.MODEL_FILE).unlink()
171 classifier = None
172 if raise_exception:
173 raise e
174 except ClassifierModelCorruptError as e:
175 # there's something wrong with the model file.
176 logger.exception(
177 "Unrecoverable error while loading document "
178 "classification model, deleting model file.",
179 )
180 Path(settings.MODEL_FILE).unlink()
181 classifier = None
182 if raise_exception:
183 raise e
184 except OSError as e:
185 logger.exception("IO error while loading document classification model")
186 classifier = None
187 if raise_exception:
188 raise e
189 except Exception as e: # pragma: no cover
190 logger.exception("Unknown error while loading document classification model")
191 classifier = None
192 if raise_exception:
193 raise e
195 return classifier
198class DocumentClassifier:
199 # v7 - Updated scikit-learn package version
200 # v8 - Added storage path classifier
201 # v9 - Changed from hashing to time/ids for re-train check
202 # v10 - HMAC-signed model file
203 # v11 - Use sample_weight for balanced training; predict_proba with threshold;
204 # drop training-only MLP state before saving
205 # Tantivy text preprocessing
206 FORMAT_VERSION = 11
208 HMAC_SIZE = 32 # SHA-256 digest length
210 def __init__(self) -> None:
211 # last time a document changed and therefore training might be required
212 self.last_doc_change_time: datetime | None = None
213 # Hash of primary keys of AUTO matching values last used in training
214 self.last_auto_type_hash: bytes | None = None
216 self.data_vectorizer = None
217 self.data_vectorizer_hash = None
218 self.tags_binarizer = None
219 self.tags_classifier = None
220 self.correspondent_classifier = None
221 self.document_type_classifier = None
222 self.storage_path_classifier = None
224 def _update_data_vectorizer_hash(self) -> None:
225 self.data_vectorizer_hash = sha256(
226 pickle.dumps(self.data_vectorizer),
227 ).hexdigest()
229 @staticmethod
230 def _new_hmac() -> hmac.HMAC:
231 return hmac.new(settings.SECRET_KEY.encode(), digestmod=sha256)
233 @staticmethod
234 def _strip_training_state(classifier: MLPClassifier) -> None:
235 """
236 Drop MLPClassifier state which is only used during fit(), never by predict().
238 The Adam optimizer keeps two moment arrays the size of the weights, and
239 without early_stopping _best_coefs/_best_intercepts are just a copy of the
240 initial random weights. Together that is 3x the size of the weights, which
241 would otherwise be pickled and loaded along with the model.
242 """
243 del classifier._optimizer
244 del classifier._best_coefs
245 del classifier._best_intercepts
247 @staticmethod
248 def _compute_hmac(data: bytes | memoryview) -> bytes:
249 mac = DocumentClassifier._new_hmac()
250 mac.update(data)
251 return mac.digest()
253 def load(self) -> None:
254 from sklearn.exceptions import InconsistentVersionWarning
256 raw = Path(settings.MODEL_FILE).read_bytes()
258 if len(raw) <= self.HMAC_SIZE:
259 raise ClassifierModelCorruptError
261 # Slice through a memoryview so the (potentially multi-GB) payload is
262 # not copied; hmac and pickle both accept buffers directly.
263 # The whole file is still verified from memory before unpickling, rather
264 # than streamed from disk, so it cannot change between check and load.
265 view = memoryview(raw)
266 signature = view[: self.HMAC_SIZE]
267 data = view[self.HMAC_SIZE :]
269 if not hmac.compare_digest(signature, self._compute_hmac(data)):
270 raise ClassifierModelCorruptError
272 # Catch warnings for processing
273 with warnings.catch_warnings(record=True) as w:
274 try:
275 (
276 schema_version,
277 self.last_doc_change_time,
278 self.last_auto_type_hash,
279 self.data_vectorizer,
280 self.tags_binarizer,
281 self.tags_classifier,
282 self.correspondent_classifier,
283 self.document_type_classifier,
284 self.storage_path_classifier,
285 ) = pickle.loads(data)
286 except Exception as err:
287 raise ClassifierModelCorruptError from err
289 if schema_version != self.FORMAT_VERSION:
290 raise IncompatibleClassifierVersionError(
291 "Cannot load classifier, incompatible versions.",
292 )
294 self._update_data_vectorizer_hash()
296 # Check for the warning about unpickling from differing versions
297 # and consider it incompatible
298 sk_learn_warning_url = (
299 "https://scikit-learn.org/stable/"
300 "model_persistence.html"
301 "#security-maintainability-limitations"
302 )
303 for warning in w:
304 # The warning is inconsistent, the MLPClassifier is a specific warning, others have not updated yet
305 if issubclass(warning.category, InconsistentVersionWarning) or (
306 issubclass(warning.category, UserWarning)
307 and sk_learn_warning_url in str(warning.message)
308 ):
309 raise IncompatibleClassifierVersionError("sklearn version update")
311 def save(self) -> None:
312 # Stream to disk instead of building the payload in memory. Protocol 5+
313 # pickles numpy arrays without copying them (the default is 4 before 3.14).
314 with _SignedFileWriter(settings.MODEL_FILE, self._new_hmac()) as f:
315 pickle.dump(
316 (
317 self.FORMAT_VERSION,
318 self.last_doc_change_time,
319 self.last_auto_type_hash,
320 self.data_vectorizer,
321 self.tags_binarizer,
322 self.tags_classifier,
323 self.correspondent_classifier,
324 self.document_type_classifier,
325 self.storage_path_classifier,
326 ),
327 f,
328 protocol=pickle.HIGHEST_PROTOCOL,
329 )
331 def train(
332 self,
333 status_callback: Callable[[str], None] | None = None,
334 ) -> bool:
335 notify = status_callback if status_callback is not None else lambda _: None
337 # Get non-inbox documents
338 docs_queryset = Document.objects.exclude(
339 tags__is_inbox_tag=True,
340 ).order_by("pk")
342 # No documents exit to train against
343 doc_count = docs_queryset.count()
344 if doc_count == 0:
345 raise ValueError("No training data available.")
347 labels_tags = []
348 labels_correspondent = []
349 labels_document_type = []
350 labels_storage_path = []
351 # Content is fetched separately later, for exactly these documents in this
352 # order, so it never all has to be in memory at once
353 doc_pks: list[int] = []
354 latest_doc_change: datetime | None = None
356 # Step 1: Extract and preprocess training data from the database.
357 logger.debug("Gathering data from database...")
358 notify(f"Gathering data from {doc_count} document(s)...")
359 hasher = sha256()
360 for doc in (
361 docs_queryset.defer("content")
362 .select_related("document_type", "correspondent", "storage_path")
363 .prefetch_related(
364 Prefetch(
365 "tags",
366 queryset=Tag.objects.filter(
367 matching_algorithm=MatchingModel.MATCH_AUTO,
368 )
369 .order_by("pk")
370 .only("pk"),
371 to_attr="auto_tags",
372 ),
373 )
374 .iterator(chunk_size=2000)
375 ):
376 doc_pks.append(doc.pk)
377 if latest_doc_change is None or doc.modified > latest_doc_change:
378 latest_doc_change = doc.modified
380 y = -1
381 dt = doc.document_type
382 if dt and dt.matching_algorithm == MatchingModel.MATCH_AUTO:
383 y = dt.pk
384 hasher.update(y.to_bytes(4, "little", signed=True))
385 labels_document_type.append(y)
387 y = -1
388 cor = doc.correspondent
389 if cor and cor.matching_algorithm == MatchingModel.MATCH_AUTO:
390 y = cor.pk
391 hasher.update(y.to_bytes(4, "little", signed=True))
392 labels_correspondent.append(y)
394 tags: list[int] = [tag.pk for tag in doc.auto_tags]
395 for tag in tags:
396 hasher.update(tag.to_bytes(4, "little", signed=True))
397 labels_tags.append(tags)
399 y = -1
400 sp = doc.storage_path
401 if sp and sp.matching_algorithm == MatchingModel.MATCH_AUTO:
402 y = sp.pk
403 hasher.update(y.to_bytes(4, "little", signed=True))
404 labels_storage_path.append(y)
406 labels_tags_unique = {tag for tags in labels_tags for tag in tags}
408 num_tags = len(labels_tags_unique)
410 # Check if retraining is actually required.
411 # A document has been updated since the classifier was trained
412 # New auto tags, types, correspondent, storage paths exist
413 if (
414 self.last_doc_change_time is not None
415 and self.last_doc_change_time >= latest_doc_change
416 ) and self.last_auto_type_hash == hasher.digest():
417 logger.info("No updates since last training")
418 # Set the classifier information into the cache
419 # Caching for 50 minutes, so slightly less than the normal retrain time
420 cache.set(
421 CLASSIFIER_MODIFIED_KEY,
422 self.last_doc_change_time,
423 CACHE_50_MINUTES,
424 )
425 cache.set(CLASSIFIER_HASH_KEY, hasher.hexdigest(), CACHE_50_MINUTES)
426 cache.set(CLASSIFIER_VERSION_KEY, self.FORMAT_VERSION, CACHE_50_MINUTES)
427 return False
429 # subtract 1 since -1 (null) is also part of the classes.
431 # union with {-1} accounts for cases where all documents have
432 # correspondents and types assigned, so -1 isn't part of labels_x, which
433 # it usually is.
434 num_correspondents: int = len(set(labels_correspondent) | {-1}) - 1
435 num_document_types: int = len(set(labels_document_type) | {-1}) - 1
436 num_storage_paths: int = len(set(labels_storage_path) | {-1}) - 1
438 logger.debug(
439 f"{len(doc_pks)} documents, {num_tags} tag(s), {num_correspondents} correspondent(s), "
440 f"{num_document_types} document type(s). {num_storage_paths} storage path(s)",
441 )
443 from sklearn.feature_extraction.text import CountVectorizer
444 from sklearn.neural_network import MLPClassifier
445 from sklearn.preprocessing import LabelBinarizer
446 from sklearn.preprocessing import MultiLabelBinarizer
448 # MLPClassifier does not support class_weight directly
449 # (https://github.com/scikit-learn/scikit-learn/issues/9113), so we use
450 # compute_sample_weight to balance classes during training and prevent
451 # over-represented correspondents from dominating predictions.
452 # https://scikit-learn.org/stable/modules/generated/sklearn.utils.class_weight.compute_sample_weight.html
453 from sklearn.utils.class_weight import compute_sample_weight
455 # Step 2: vectorize data
456 logger.debug("Vectorizing data...")
457 notify("Vectorizing document content...")
459 def content_generator() -> Iterator[str]:
460 """
461 Generates the content for documents, in the same order as the labels,
462 fetching it a chunk at a time
463 """
464 for start in range(0, len(doc_pks), _CONTENT_CHUNK_SIZE):
465 chunk = doc_pks[start : start + _CONTENT_CHUNK_SIZE]
466 docs = Document.objects.only("content").order_by().in_bulk(chunk)
467 for pk in chunk:
468 # A document deleted since its labels were gathered still
469 # needs a row, so labels and content stay aligned
470 doc = docs.get(pk)
471 yield self.preprocess_content(
472 doc.content if doc is not None else "",
473 )
475 self.data_vectorizer = CountVectorizer(
476 analyzer="word",
477 ngram_range=(1, 2),
478 min_df=0.01,
479 )
481 data_vectorized: ndarray = self.data_vectorizer.fit_transform(
482 content_generator(),
483 )
485 # See the notes here:
486 # https://scikit-learn.org/stable/modules/generated/sklearn.feature_extraction.text.CountVectorizer.html
487 # This attribute isn't needed to function and can be large
488 self.data_vectorizer.stop_words_ = None
490 # Step 3: train the classifiers
491 if num_tags > 0:
492 logger.debug("Training tags classifier...")
493 notify(f"Training tags classifier ({num_tags} tag(s))...")
495 if num_tags == 1:
496 # Special case where only one tag has auto:
497 # Fallback to binary classification.
498 labels_tags = [
499 label[0] if len(label) == 1 else -1 for label in labels_tags
500 ]
501 self.tags_binarizer = LabelBinarizer()
502 labels_tags_vectorized: ndarray = self.tags_binarizer.fit_transform(
503 labels_tags,
504 ).ravel()
505 else:
506 self.tags_binarizer = MultiLabelBinarizer()
507 labels_tags_vectorized = self.tags_binarizer.fit_transform(labels_tags)
509 self.tags_classifier = MLPClassifier(tol=0.01, random_state=0)
510 self.tags_classifier.fit(data_vectorized, labels_tags_vectorized)
511 self._strip_training_state(self.tags_classifier)
512 else:
513 self.tags_classifier = None
514 logger.debug("There are no tags. Not training tags classifier.")
516 if num_correspondents > 0:
517 logger.debug("Training correspondent classifier...")
518 notify(
519 f"Training correspondent classifier ({num_correspondents} correspondent(s))...",
520 )
521 self.correspondent_classifier = MLPClassifier(tol=0.01, random_state=0)
522 self.correspondent_classifier.fit(
523 data_vectorized,
524 labels_correspondent,
525 sample_weight=compute_sample_weight("balanced", labels_correspondent),
526 )
527 self._strip_training_state(self.correspondent_classifier)
528 else:
529 self.correspondent_classifier = None
530 logger.debug(
531 "There are no correspondents. Not training correspondent classifier.",
532 )
534 if num_document_types > 0:
535 logger.debug("Training document type classifier...")
536 notify(
537 f"Training document type classifier ({num_document_types} type(s))...",
538 )
539 self.document_type_classifier = MLPClassifier(tol=0.01, random_state=0)
540 self.document_type_classifier.fit(
541 data_vectorized,
542 labels_document_type,
543 sample_weight=compute_sample_weight("balanced", labels_document_type),
544 )
545 self._strip_training_state(self.document_type_classifier)
546 else:
547 self.document_type_classifier = None
548 logger.debug(
549 "There are no document types. Not training document type classifier.",
550 )
552 if num_storage_paths > 0:
553 logger.debug(
554 "Training storage paths classifier...",
555 )
556 notify(f"Training storage path classifier ({num_storage_paths} path(s))...")
557 self.storage_path_classifier = MLPClassifier(tol=0.01, random_state=0)
558 self.storage_path_classifier.fit(
559 data_vectorized,
560 labels_storage_path,
561 sample_weight=compute_sample_weight("balanced", labels_storage_path),
562 )
563 self._strip_training_state(self.storage_path_classifier)
564 else:
565 self.storage_path_classifier = None
566 logger.debug(
567 "There are no storage paths. Not training storage path classifier.",
568 )
570 self.last_doc_change_time = latest_doc_change
571 self.last_auto_type_hash = hasher.digest()
572 self._update_data_vectorizer_hash()
574 # Set the classifier information into the cache
575 # Caching for 50 minutes, so slightly less than the normal retrain time
576 cache.set(CLASSIFIER_MODIFIED_KEY, self.last_doc_change_time, CACHE_50_MINUTES)
577 cache.set(CLASSIFIER_HASH_KEY, hasher.hexdigest(), CACHE_50_MINUTES)
578 cache.set(CLASSIFIER_VERSION_KEY, self.FORMAT_VERSION, CACHE_50_MINUTES)
580 return True
582 def preprocess_content(self, content: str) -> str:
583 """
584 Process the contents of a document, distilling it down into
585 words which are meaningful to the content.
586 """
587 language = settings.CLASSIFIER_LANGUAGE
588 content = unicodedata.normalize("NFC", content)
589 if language is None:
590 return " ".join(
591 match.group().lower() for match in RE_WORD.finditer(content)
592 )
593 return " ".join(_text_analyzer(language).analyze(content))
595 def _get_vectorizer_cache_key(self, content: str):
596 hash = sha256(content.encode())
597 hash.update(
598 f"|{self.FORMAT_VERSION}|{settings.CLASSIFIER_LANGUAGE}|{self.data_vectorizer_hash}".encode(),
599 )
600 return f"vectorized_content_{hash.hexdigest()}"
602 def _vectorize(self, content: str):
603 key = self._get_vectorizer_cache_key(content)
604 serialized_result = read_cache.get(key)
605 if serialized_result is None:
606 result = self.data_vectorizer.transform([self.preprocess_content(content)])
607 read_cache.set(key, signed_pickle_dumps(result), CACHE_5_MINUTES)
608 else:
609 try:
610 result = signed_pickle_loads(serialized_result)
611 except SignedPickleError:
612 result = self.data_vectorizer.transform(
613 [self.preprocess_content(content)],
614 )
615 read_cache.set(key, signed_pickle_dumps(result), CACHE_5_MINUTES)
616 else:
617 read_cache.touch(key, CACHE_5_MINUTES)
618 return result
620 def predict_correspondent(self, content: str) -> int | None:
621 if self.correspondent_classifier:
622 X = self._vectorize(content)
623 predicted_id = _predict_with_threshold(
624 self.correspondent_classifier,
625 X,
626 settings.CLASSIFIER_MATCH_THRESHOLD,
627 )
628 return predicted_id
629 return None
631 def predict_document_type(self, content: str) -> int | None:
632 if self.document_type_classifier:
633 X = self._vectorize(content)
634 predicted_id = _predict_with_threshold(
635 self.document_type_classifier,
636 X,
637 settings.CLASSIFIER_MATCH_THRESHOLD,
638 )
639 return predicted_id
640 return None
642 def predict_tags(self, content: str) -> list[int]:
643 from sklearn.utils.multiclass import type_of_target
645 if self.tags_classifier:
646 X = self._vectorize(content)
647 y = self.tags_classifier.predict(X)
648 tags_ids = self.tags_binarizer.inverse_transform(y)[0]
649 if type_of_target(y).startswith("multilabel"):
650 # the usual case when there are multiple tags.
651 return list(tags_ids)
652 elif type_of_target(y) == "binary" and tags_ids != -1:
653 # This is for when we have binary classification with only one
654 # tag and the result is to assign this tag.
655 return [tags_ids]
656 else:
657 # Usually binary as well with -1 as the result, but we're
658 # going to catch everything else here as well.
659 return []
660 else:
661 return []
663 def predict_storage_path(self, content: str) -> int | None:
664 if self.storage_path_classifier:
665 X = self._vectorize(content)
666 predicted_id = _predict_with_threshold(
667 self.storage_path_classifier,
668 X,
669 settings.CLASSIFIER_MATCH_THRESHOLD,
670 )
671 return predicted_id
672 return None