From 42da9317b2bdad37fa08c7f0984a494ab65edf89 Mon Sep 17 00:00:00 2001 From: J08nY Date: Fri, 10 Feb 2023 20:09:12 +0100 Subject: Drop inheritance from sklearn objects. It is unused and causes big memory usag spikes on their import due to their inefficient import practice: https://github.com/scikit-learn/scikit-learn/issues/25590 We don't actually use their API anywhere (afaik). --- src/sec_certs/model/cpe_matching.py | 5 ++--- src/sec_certs/model/sar_transformer.py | 9 +++++---- 2 files changed, 7 insertions(+), 7 deletions(-) (limited to 'src') diff --git a/src/sec_certs/model/cpe_matching.py b/src/sec_certs/model/cpe_matching.py index 2c6b3f07..5d08d7af 100644 --- a/src/sec_certs/model/cpe_matching.py +++ b/src/sec_certs/model/cpe_matching.py @@ -8,7 +8,6 @@ from typing import Pattern import spacy from rapidfuzz import fuzz -from sklearn.base import BaseEstimator from sec_certs import cert_rules, constants from sec_certs.sample.cpe import CPE @@ -17,10 +16,10 @@ from sec_certs.utils.tqdm import tqdm logger = logging.getLogger(__name__) -class CPEClassifier(BaseEstimator): +class CPEClassifier: """ Class that can predict CPE matches for certificate instances. - Adheres to sklearn BaseEstimator interface. + Adheres to sklearn `sklearn.base.BaseEstimator` interface. Fit method is called on list of CPEs and build two look-up dictionaries, see description of attributes. """ diff --git a/src/sec_certs/model/sar_transformer.py b/src/sec_certs/model/sar_transformer.py index 20ce9dec..a60f7495 100644 --- a/src/sec_certs/model/sar_transformer.py +++ b/src/sec_certs/model/sar_transformer.py @@ -3,8 +3,6 @@ from __future__ import annotations import logging from typing import Dict, Iterable, cast -from sklearn.base import BaseEstimator, TransformerMixin - from sec_certs.sample.cc import CCCertificate from sec_certs.sample.sar import SAR, SAR_DICT_KEY @@ -12,10 +10,10 @@ logger = logging.getLogger(__name__) # TODO: Right now we ignore number of ocurrences for final SAR selection. If we keep it this way, we can discard that variable -class SARTransformer(BaseEstimator, TransformerMixin): +class SARTransformer: """ Class for transforming SARs defined in st_keywords and report_keywords dictionaries into SAR objects. - This class implements sklearn transformer interface, so fit_transform() can be called on it. + This class implements `sklearn.base.Transformer` interface, so fit_transform() can be called on it. """ def fit(self, certificates: Iterable[CCCertificate]) -> SARTransformer: @@ -27,6 +25,9 @@ class SARTransformer(BaseEstimator, TransformerMixin): """ return self + def fit_transform(self, X, y=None, **fit_params): + return self.fit(X).transform(X) + def transform(self, certificates: Iterable[CCCertificate]) -> list[set[SAR] | None]: """ Just a wrapper around transform_single_cert() called on an iterable of CCCertificate. -- cgit v1.3.1