Coverage for documents/classifier.py: 16%

336 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 09:07 +0000

1from __future__ import annotations 

2 

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 

13 

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 

21 

22 import tantivy 

23 from numpy import ndarray 

24 from sklearn.neural_network import MLPClassifier 

25 

26from django.conf import settings 

27from django.core.cache import cache 

28from django.core.cache import caches 

29from django.db.models import Prefetch 

30 

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 

43 

44logger = logging.getLogger("paperless.classifier") 

45 

46 

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. 

52 

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]) 

59 

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 

65 

66 

67read_cache = caches["read-cache"] 

68 

69 

70RE_WORD = re.compile(r"\b[\w]+\b") # words that may contain digits 

71 

72# Documents whose content is fetched per query while training 

73_CONTENT_CHUNK_SIZE = 1000 

74 

75 

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. 

80 

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 """ 

85 

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 

91 

92 def __enter__(self) -> Self: 

93 self._file = self._temp.open("wb") 

94 self._file.write(bytes(self._mac.digest_size)) 

95 return self 

96 

97 def write(self, data: bytes | memoryview) -> int: 

98 self._mac.update(data) 

99 return self._file.write(data) 

100 

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) 

117 

118 

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. 

123 

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 

128 

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 ) 

144 

145 

146class IncompatibleClassifierVersionError(Exception): 

147 def __init__(self, message: str, *args: object) -> None: 

148 self.message: str = message 

149 super().__init__(*args) 

150 

151 

152class ClassifierModelCorruptError(Exception): 

153 pass 

154 

155 

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 

163 

164 classifier = DocumentClassifier() 

165 try: 

166 classifier.load() 

167 

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 

194 

195 return classifier 

196 

197 

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 

207 

208 HMAC_SIZE = 32 # SHA-256 digest length 

209 

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 

215 

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 

223 

224 def _update_data_vectorizer_hash(self) -> None: 

225 self.data_vectorizer_hash = sha256( 

226 pickle.dumps(self.data_vectorizer), 

227 ).hexdigest() 

228 

229 @staticmethod 

230 def _new_hmac() -> hmac.HMAC: 

231 return hmac.new(settings.SECRET_KEY.encode(), digestmod=sha256) 

232 

233 @staticmethod 

234 def _strip_training_state(classifier: MLPClassifier) -> None: 

235 """ 

236 Drop MLPClassifier state which is only used during fit(), never by predict(). 

237 

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 

246 

247 @staticmethod 

248 def _compute_hmac(data: bytes | memoryview) -> bytes: 

249 mac = DocumentClassifier._new_hmac() 

250 mac.update(data) 

251 return mac.digest() 

252 

253 def load(self) -> None: 

254 from sklearn.exceptions import InconsistentVersionWarning 

255 

256 raw = Path(settings.MODEL_FILE).read_bytes() 

257 

258 if len(raw) <= self.HMAC_SIZE: 

259 raise ClassifierModelCorruptError 

260 

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 :] 

268 

269 if not hmac.compare_digest(signature, self._compute_hmac(data)): 

270 raise ClassifierModelCorruptError 

271 

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 

288 

289 if schema_version != self.FORMAT_VERSION: 

290 raise IncompatibleClassifierVersionError( 

291 "Cannot load classifier, incompatible versions.", 

292 ) 

293 

294 self._update_data_vectorizer_hash() 

295 

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") 

310 

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 ) 

330 

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 

336 

337 # Get non-inbox documents 

338 docs_queryset = Document.objects.exclude( 

339 tags__is_inbox_tag=True, 

340 ).order_by("pk") 

341 

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.") 

346 

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 

355 

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 

379 

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) 

386 

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) 

393 

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) 

398 

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) 

405 

406 labels_tags_unique = {tag for tags in labels_tags for tag in tags} 

407 

408 num_tags = len(labels_tags_unique) 

409 

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 

428 

429 # subtract 1 since -1 (null) is also part of the classes. 

430 

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 

437 

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 ) 

442 

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 

447 

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 

454 

455 # Step 2: vectorize data 

456 logger.debug("Vectorizing data...") 

457 notify("Vectorizing document content...") 

458 

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 ) 

474 

475 self.data_vectorizer = CountVectorizer( 

476 analyzer="word", 

477 ngram_range=(1, 2), 

478 min_df=0.01, 

479 ) 

480 

481 data_vectorized: ndarray = self.data_vectorizer.fit_transform( 

482 content_generator(), 

483 ) 

484 

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 

489 

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))...") 

494 

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) 

508 

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.") 

515 

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 ) 

533 

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 ) 

551 

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 ) 

569 

570 self.last_doc_change_time = latest_doc_change 

571 self.last_auto_type_hash = hasher.digest() 

572 self._update_data_vectorizer_hash() 

573 

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) 

579 

580 return True 

581 

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)) 

594 

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()}" 

601 

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 

619 

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 

630 

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 

641 

642 def predict_tags(self, content: str) -> list[int]: 

643 from sklearn.utils.multiclass import type_of_target 

644 

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 [] 

662 

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