aboutsummaryrefslogtreecommitdiffhomepage
path: root/src
diff options
context:
space:
mode:
authorJ08nY2023-02-10 20:09:12 +0100
committerJ08nY2023-02-10 20:10:29 +0100
commit42da9317b2bdad37fa08c7f0984a494ab65edf89 (patch)
tree8a4be55473d1febb815d7fa537fc1c8c347cfa53 /src
parent410dbe6a0c6a1a609a129b3919ea58e4453cdeb1 (diff)
downloadsec-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.py5
-rw-r--r--src/sec_certs/model/sar_transformer.py9
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.