diff options
| author | J08nY | 2023-02-10 20:09:12 +0100 |
|---|---|---|
| committer | J08nY | 2023-02-10 20:10:29 +0100 |
| commit | 42da9317b2bdad37fa08c7f0984a494ab65edf89 (patch) | |
| tree | 8a4be55473d1febb815d7fa537fc1c8c347cfa53 /src | |
| parent | 410dbe6a0c6a1a609a129b3919ea58e4453cdeb1 (diff) | |
| download | sec-certs-42da9317b2bdad37fa08c7f0984a494ab65edf89.tar.gz sec-certs-42da9317b2bdad37fa08c7f0984a494ab65edf89.tar.zst sec-certs-42da9317b2bdad37fa08c7f0984a494ab65edf89.zip | |
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).
Diffstat (limited to 'src')
| -rw-r--r-- | src/sec_certs/model/cpe_matching.py | 5 | ||||
| -rw-r--r-- | src/sec_certs/model/sar_transformer.py | 9 |
2 files changed, 7 insertions, 7 deletions
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. |
