aboutsummaryrefslogtreecommitdiffhomepage
path: root/src
diff options
context:
space:
mode:
authoradamjanovsky2023-11-14 10:04:13 +0100
committeradamjanovsky2023-11-14 10:04:13 +0100
commit80190b01aeda844b9d3ea8684284130c44f1453e (patch)
tree6fbcabd9cda272b9a5d64c8e61c7d3b914351f93 /src
parent9cdf4801f93243e682b43be0a52956c0f9fad377 (diff)
downloadsec-certs-80190b01aeda844b9d3ea8684284130c44f1453e.tar.gz
sec-certs-80190b01aeda844b9d3ea8684284130c44f1453e.tar.zst
sec-certs-80190b01aeda844b9d3ea8684284130c44f1453e.zip
bump references
Diffstat (limited to 'src')
-rw-r--r--src/sec_certs/constants.py7
-rw-r--r--src/sec_certs/data/reference_annotations/readme.md31
-rw-r--r--src/sec_certs/dataset/cc.py69
-rw-r--r--src/sec_certs/dataset/fips.py2
-rw-r--r--src/sec_certs/model/__init__.py5
-rw-r--r--src/sec_certs/model/reference_finder.py (renamed from src/sec_certs/model/references/reference_finder.py)0
-rw-r--r--src/sec_certs/model/references_nlp/__init__.py13
-rw-r--r--src/sec_certs/model/references_nlp/annotator.py (renamed from src/sec_certs/model/references/annotator.py)25
-rw-r--r--src/sec_certs/model/references_nlp/annotator_trainer.py (renamed from src/sec_certs/model/references/annotator_trainer.py)45
-rw-r--r--src/sec_certs/model/references_nlp/evaluation.py90
-rw-r--r--src/sec_certs/model/references_nlp/feature_extraction.py542
-rw-r--r--src/sec_certs/model/references_nlp/segment_extractor.py (renamed from src/sec_certs/model/references/segment_extractor.py)52
-rw-r--r--src/sec_certs/model/references_nlp/training.py68
13 files changed, 759 insertions, 190 deletions
diff --git a/src/sec_certs/constants.py b/src/sec_certs/constants.py
index b28b925d..773a5720 100644
--- a/src/sec_certs/constants.py
+++ b/src/sec_certs/constants.py
@@ -1,6 +1,11 @@
import re
from pathlib import Path
-from typing import Final
+from typing import Final, Literal
+
+RANDOM_STATE: Final[int] = 42
+REF_ANNOTATION_MODES = Literal["training", "evaluation", "production", "cross-validation"]
+REF_EMBEDDING_METHOD = Literal["tf_idf", "transformer"]
+
DUMMY_NONEXISTING_PATH = Path("/this/is/dummy/nonexisting/path")
diff --git a/src/sec_certs/data/reference_annotations/readme.md b/src/sec_certs/data/reference_annotations/readme.md
index b10a2e30..8521eead 100644
--- a/src/sec_certs/data/reference_annotations/readme.md
+++ b/src/sec_certs/data/reference_annotations/readme.md
@@ -54,35 +54,8 @@ These can be further merged into the following super-categories:
The inter-annotator agreement is measured both with Cohen's Kappa and with percentage. The results are as follows:
| Cohen's Kappa | Percentage |
-|---------------|------------|
+| ------------- | ---------- |
| 0.71 | 0.82 |
-The code used to measure the agreement is:
+The code used to measure the agreement is stored in `notebooks/cc/reference_annotations/inter_annotator_agreement.ipynb`.
-```python
-import pandas as pd
-from pathlib import Path
-from sklearn.metrics import cohen_kappa_score
-
-def load_all_dataframes(base_folder: Path) -> pd.DataFrame:
- splits = ["train", "valid", "test"]
-
- df_train, df_valid, df_test = pd.DataFrame(), pd.DataFrame(), pd.DataFrame()
- for split in splits:
- df = pd.read_csv(base_folder / f"{split}.csv")
- if split == "train":
- df_train = df
- elif split == "valid":
- df_valid = df
- else:
- df_test = df
-
- return pd.concat([df_train, df_valid, df_test])
-
-adam_df = load_all_dataframes(Path("./src/sec_certs/data/reference_annotations/adam"))
-jano_df = load_all_dataframes(Path("./src/sec_certs/data/reference_annotations/jano"))
-agreement_series = adam_df.label == jano_df.label
-
-print(f"Cohen's Kappa: {cohen_kappa_score(adam_df.label, jano_df.label)}")
-print(f"Percentage agreement: {agreement_series.loc[agreement_series == True].count() / agreement_series.count()}")
-```
diff --git a/src/sec_certs/dataset/cc.py b/src/sec_certs/dataset/cc.py
index ddd633e7..afde8bd2 100644
--- a/src/sec_certs/dataset/cc.py
+++ b/src/sec_certs/dataset/cc.py
@@ -7,7 +7,7 @@ import tempfile
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
-from typing import ClassVar, Iterator, Literal, cast
+from typing import ClassVar, Iterator, cast
import numpy as np
import pandas as pd
@@ -22,14 +22,11 @@ from sec_certs.dataset.cve import CVEDataset
from sec_certs.dataset.dataset import AuxiliaryDatasets, Dataset, logger
from sec_certs.dataset.protection_profile import ProtectionProfileDataset
from sec_certs.model import (
- ReferenceAnnotator,
- ReferenceAnnotatorTrainer,
ReferenceFinder,
SARTransformer,
TransitiveVulnerabilityFinder,
)
from sec_certs.model.cc_matching import CCSchemeMatcher
-from sec_certs.model.references.segment_extractor import ReferenceSegmentExtractor
from sec_certs.sample.cc import CCCertificate
from sec_certs.sample.cc_certificate_id import CertificateId
from sec_certs.sample.cc_maintenance_update import CCMaintenanceUpdate
@@ -38,7 +35,6 @@ from sec_certs.sample.protection_profile import ProtectionProfile
from sec_certs.serialization.json import ComplexSerializableType, serialize
from sec_certs.utils import helpers
from sec_certs.utils import parallel_processing as cert_processing
-from sec_certs.utils.nlp import prec_recall_metric
@dataclass
@@ -729,7 +725,6 @@ class CCDataset(Dataset[CCCertificate, CCAuxiliaryDatasets], ComplexSerializable
self._compute_scheme_data()
self._compute_cert_labs()
self._compute_sars()
- self.annotate_references()
def _compute_sars(self) -> None:
logger.info("Computing heuristics: Computing SARs")
@@ -838,68 +833,6 @@ class CCDataset(Dataset[CCCertificate, CCAuxiliaryDatasets], ComplexSerializable
return update_dset
- def annotate_references(self, mode: Literal["training", "production"] = "production"):
- """
- Fills in `cert.heuristics.annotated_references` with reference labels.
- This requires (a) a pre-trained `ReferenceAnnotator` or (b) to train a `ReferenceAnnotator`.
- The behaviour is controlled by config keys `config.cc_reference_annotator_should_train` and
- `config.cc_reference_annotator_path`. If should_train is False and path is provided, a pre-trained annotator
- will be loaded and used.
- """
- if not self.state.pdfs_converted:
- logger.info(
- "Attempting run analysis of txt files while not having the pdf->txt conversion done. Returning."
- )
- return
- if not self.state.auxiliary_datasets_processed:
- logger.info(
- "Attempting to run analysis of certifies while not having the auxiliary datasets processed. Returning."
- )
-
- if not config.cc_reference_annotator_should_train:
- model_dir = (
- config.cc_reference_annotator_dir if config.cc_reference_annotator_dir else self.reference_annotator_dir
- )
- try:
- annotator = ReferenceAnnotator.from_pretrained(model_dir)
- except Exception:
- logger.error(
- "annotate_references() method was called with `config.cc_reference_annotator_should_train=False`."
- f"Further, the model was not found either at `config.cc_reference_annotator_dir={config.cc_reference_annotator_dir}`"
- f"nor at {self.reference_annotator_dir}. Either: (a) allow training with `config.cc_reference_annotator_should_train=True`;"
- "(b) set path to model with `config.cc_reference_annotator_path`; (c) paste the model into {self.reference_annotator_dir}. Returning."
- )
- return
-
- logger.info("Extracting segments of text relevant for reference annotations.")
- df = ReferenceSegmentExtractor()(self.certs.values())
- if config.cc_reference_annotator_should_train:
- annotator = self._train_reference_annotator(df, mode=mode)
-
- logger.info("Predicting reference labels")
- df = annotator.predict_df(df)
- refs: dict[str, dict[str, str]] = dict.fromkeys(df.dgst, {})
- for dgst, cert_id, label in zip(df.dgst, df.referenced_cert_id, df.y_pred):
- refs[dgst][cert_id] = label
-
- for dgst, value in refs.items():
- self[dgst].heuristics.annotated_references = value
-
- def _train_reference_annotator(
- self, df: pd.DataFrame, save_model: bool = True, mode: Literal["training", "production"] = "production"
- ) -> ReferenceAnnotator:
- trainer = ReferenceAnnotatorTrainer.from_df(df, prec_recall_metric, "transformer", mode)
- logger.info(
- "Training ReferenceAnnotator on {df.shape[0]} samples ({df.loc[df.split == 'train'].shape[0]}/{df.loc[df.split == 'valid'].shape[0]}/{df.loc[df.split == 'test'].shape[0]}) (train/valid/test)."
- )
- trainer.train()
- logger.info(trainer.evaluate())
-
- if save_model:
- trainer.clf.save_pretrained(self.reference_annotator_dir)
-
- return trainer.clf
-
def process_schemes(self, to_download: bool = True, only_schemes: set[str] | None = None) -> CCSchemeDataset:
"""
Downloads or loads from json a dataset of CC scheme data.
diff --git a/src/sec_certs/dataset/fips.py b/src/sec_certs/dataset/fips.py
index f5dda10d..24536b9b 100644
--- a/src/sec_certs/dataset/fips.py
+++ b/src/sec_certs/dataset/fips.py
@@ -17,7 +17,7 @@ from sec_certs.dataset.cpe import CPEDataset
from sec_certs.dataset.cve import CVEDataset
from sec_certs.dataset.dataset import AuxiliaryDatasets, Dataset
from sec_certs.dataset.fips_algorithm import FIPSAlgorithmDataset
-from sec_certs.model.references.reference_finder import ReferenceFinder
+from sec_certs.model.reference_finder import ReferenceFinder
from sec_certs.model.transitive_vulnerability_finder import TransitiveVulnerabilityFinder
from sec_certs.sample.fips import FIPSCertificate
from sec_certs.serialization.json import ComplexSerializableType, serialize
diff --git a/src/sec_certs/model/__init__.py b/src/sec_certs/model/__init__.py
index 4a3a14c0..69cbe194 100644
--- a/src/sec_certs/model/__init__.py
+++ b/src/sec_certs/model/__init__.py
@@ -6,10 +6,7 @@ leveraged by members of Dataset package and are directly applied on members of S
from sec_certs.model.cc_matching import CCSchemeMatcher
from sec_certs.model.cpe_matching import CPEClassifier
from sec_certs.model.fips_matching import FIPSProcessMatcher
-from sec_certs.model.references.annotator import ReferenceAnnotator
-from sec_certs.model.references.annotator_trainer import ReferenceAnnotatorTrainer
-from sec_certs.model.references.reference_finder import ReferenceFinder
-from sec_certs.model.references.segment_extractor import ReferenceSegmentExtractor
+from sec_certs.model.reference_finder import ReferenceFinder
from sec_certs.model.sar_transformer import SARTransformer
from sec_certs.model.transitive_vulnerability_finder import TransitiveVulnerabilityFinder
diff --git a/src/sec_certs/model/references/reference_finder.py b/src/sec_certs/model/reference_finder.py
index 94a3b29f..94a3b29f 100644
--- a/src/sec_certs/model/references/reference_finder.py
+++ b/src/sec_certs/model/reference_finder.py
diff --git a/src/sec_certs/model/references_nlp/__init__.py b/src/sec_certs/model/references_nlp/__init__.py
new file mode 100644
index 00000000..82c95aa1
--- /dev/null
+++ b/src/sec_certs/model/references_nlp/__init__.py
@@ -0,0 +1,13 @@
+# ruff: noqa: F401
+try:
+ import catboost
+ import optuna
+ import plotly.express
+ import setfit
+ import sklearn
+ import umap
+except ImportError as e:
+ print(e)
+ print(
+ f"Requirements for ML annotation of references not met. Please run `pip install sec-certs[nlp]` or install `pip install -r requirements/nlp_requirements.txt."
+ )
diff --git a/src/sec_certs/model/references/annotator.py b/src/sec_certs/model/references_nlp/annotator.py
index 76676895..1cc816f5 100644
--- a/src/sec_certs/model/references/annotator.py
+++ b/src/sec_certs/model/references_nlp/annotator.py
@@ -2,7 +2,6 @@ from __future__ import annotations
import json
import logging
-import re
from collections import Counter
from dataclasses import dataclass
from pathlib import Path
@@ -27,9 +26,7 @@ class ReferenceAnnotator:
_model: Any
_label_mapping: dict[int, str]
_soft_voting_power: int = 2
- _use_analytical_rule_name_similarity: bool = True
- # TODO: This does not load hyperparameters, only the model and label mapping
@classmethod
def from_pretrained(cls, model_dir: str | Path) -> ReferenceAnnotator:
"""
@@ -48,7 +45,6 @@ class ReferenceAnnotator:
return cls(model, label_mapping)
- # TODO: This does not save hyperparameters, only the model and label mapping
def save_pretrained(self, model_dir: str | Path):
"""
Will dump _model and _label_mapping into a directory.
@@ -95,31 +91,10 @@ class ReferenceAnnotator:
"""
WIll read df.segments and populate the dataframe with predictions.
"""
-
- def matches_reevaluation(segments: list[str]) -> bool:
- regex_a = r"This is a re-?\s?certification based on (the\s){1,2}referenced product"
- regex_b = r"Re-?\s?Zertifizierung basierend auf (the\s){1,2}referenced product"
- return any(
- re.search(regex_a, segment, re.IGNORECASE) or re.search(regex_b, segment, re.IGNORECASE)
- for segment in segments
- )
-
df_new = df.copy()
y_proba = self.predict_proba(df_new.segments)
df_new["y_proba"] = y_proba
df_new["y_pred"] = self.predict(df_new.segments)
-
- if self._use_analytical_rule_name_similarity:
- df_new.loc[
- (df_new.name_similarity == 100)
- & (df_new.name_len_diff < 5)
- & ((df_new.y_pred != "RE-EVALUATION") & (df_new.y_pred != "PREVIOUS_VERSION")),
- ["y_pred"],
- ] = "PREVIOUS_VERSION"
-
- df_new["maches_reevaluation"] = df_new.segments.map(matches_reevaluation)
- df_new.loc[df_new.maches_reevaluation, ["y_pred"]] = "RE-EVALUATION"
-
df_new["correct"] = df_new.apply(
lambda row: row["y_pred"] == row["label"] if not pd.isnull(row["label"]) else np.NaN, axis=1
)
diff --git a/src/sec_certs/model/references/annotator_trainer.py b/src/sec_certs/model/references_nlp/annotator_trainer.py
index fcc08422..1f973842 100644
--- a/src/sec_certs/model/references/annotator_trainer.py
+++ b/src/sec_certs/model/references_nlp/annotator_trainer.py
@@ -2,47 +2,57 @@ from __future__ import annotations
import logging
from functools import partial
-from typing import Callable, Literal
+from typing import Callable, Final, Literal
import pandas as pd
from datasets import ClassLabel, Dataset, Features, NamedSplit, Value
from sentence_transformers.losses import CosineSimilarityLoss
from setfit import SetFitModel, SetFitTrainer
-from sklearn.metrics import f1_score
+from sklearn.metrics import accuracy_score, balanced_accuracy_score, f1_score
-from sec_certs.model.references.annotator import ReferenceAnnotator
+from sec_certs.model.references_nlp.annotator import ReferenceAnnotator
from sec_certs.utils.nlp import prepare_reference_annotations_df
logger = logging.getLogger(__name__)
class ReferenceAnnotatorTrainer:
+ METRIC_TO_USE: Final[dict[str, Callable]] = {
+ "accuracy": accuracy_score,
+ "balanced_accuracy": balanced_accuracy_score,
+ "f1": partial(f1_score, average="weighted", zero_division=0),
+ }
+
def __init__(
self,
train_dataset: pd.DataFrame,
eval_dataset: pd.DataFrame,
metric: Callable,
- use_analytical_rule_name_similarity: bool = True,
n_iterations: int = 20,
+ learning_rate: float = 2e-5,
n_epochs: int = 1,
batch_size: int = 16,
- segmenter_metric: Literal["accuracy", "f1"] = "accuracy",
+ segmenter_metric: Literal["accuracy", "f1", "balanced_accuracy"] = "accuracy",
ensemble_soft_voting_power: int = 2,
+ show_progress_bar: bool = True,
):
self._train_dataset = train_dataset
self._eval_dataset = eval_dataset
self._metric = metric
- self.use_analytical_rule_name_similarity = use_analytical_rule_name_similarity
self.n_iterations = n_iterations
+ self.learning_rate = learning_rate
self.n_epochs = n_epochs
self.batch_size = batch_size
self.segmenter_metric = segmenter_metric
self.ensemble_soft_voting_power = ensemble_soft_voting_power
+ self.show_progress_bar = show_progress_bar
self._model, self._trainer, self.label_mapping = self._init_trainer()
self.clf = ReferenceAnnotator(
- self._model, self.label_mapping, self.ensemble_soft_voting_power, self.use_analytical_rule_name_similarity
+ self._model,
+ self.label_mapping,
+ self.ensemble_soft_voting_power,
)
@classmethod
@@ -50,19 +60,21 @@ class ReferenceAnnotatorTrainer:
cls,
df: pd.DataFrame,
metric: Callable,
- mode: Literal["training", "evaluation", "production"] = "training",
- use_analytical_rule_name_similarity: bool = True,
+ mode: Literal["training", "evaluation", "production", "cross-validation"] = "training",
n_iterations: int = 20,
+ learning_rate: float = 2e-5,
n_epochs: int = 1,
batch_size: int = 16,
- segmenter_metric: Literal["accuracy", "f1"] = "accuracy",
+ segmenter_metric: Literal["accuracy", "f1", "balanced_accuracy"] = "accuracy",
ensemble_soft_voting_power: int = 2,
+ show_progress_bar: bool = True,
):
df = prepare_reference_annotations_df(df)
dataset_generation_method = {
"training": ReferenceAnnotatorTrainer.split_df_for_training,
"evaluation": ReferenceAnnotatorTrainer.split_df_for_evaluation,
"production": ReferenceAnnotatorTrainer.split_df_for_production,
+ "cross-validation": ReferenceAnnotatorTrainer.split_df_for_training,
}
train_dataset, eval_dataset = dataset_generation_method[mode](df)
@@ -70,12 +82,13 @@ class ReferenceAnnotatorTrainer:
train_dataset,
eval_dataset,
metric,
- use_analytical_rule_name_similarity,
n_iterations,
+ learning_rate,
n_epochs,
batch_size,
segmenter_metric,
ensemble_soft_voting_power,
+ show_progress_bar,
)
@staticmethod
@@ -109,17 +122,13 @@ class ReferenceAnnotatorTrainer:
internal_train_dataset = internal_train_dataset.align_labels_with_mapping(label2id, "label")
internal_validation_dataset = internal_validation_dataset.align_labels_with_mapping(label2id, "label")
- if self.segmenter_metric == "accuracy":
- metric_to_use = "accuracy"
- else:
- metric_to_use = partial(f1_score, average="weighted", zero_division=0)
-
trainer = SetFitTrainer(
model=model,
train_dataset=internal_train_dataset,
eval_dataset=internal_validation_dataset,
loss_class=CosineSimilarityLoss,
- metric=metric_to_use,
+ metric=self.METRIC_TO_USE[self.segmenter_metric],
+ learning_rate=self.learning_rate,
batch_size=self.batch_size,
num_iterations=self.n_iterations, # The number of text pairs to generate for contrastive learning
num_epochs=self.n_epochs, # The number of epochs to use for contrastive learning
@@ -144,7 +153,7 @@ class ReferenceAnnotatorTrainer:
return Dataset.from_pandas(df_to_use, features=features, split=split, preserve_index=False)
def train(self):
- self._trainer.train(show_progress_bar=True)
+ self._trainer.train(show_progress_bar=self.show_progress_bar)
def evaluate(self):
print("Internal evaluation (of model working on individual segments)")
diff --git a/src/sec_certs/model/references_nlp/evaluation.py b/src/sec_certs/model/references_nlp/evaluation.py
new file mode 100644
index 00000000..dc9b4b13
--- /dev/null
+++ b/src/sec_certs/model/references_nlp/evaluation.py
@@ -0,0 +1,90 @@
+import logging
+from pathlib import Path
+from typing import Literal
+
+import matplotlib.pyplot as plt
+import numpy as np
+import pandas as pd
+import plotly.express as px
+from catboost import CatBoostClassifier
+from sklearn.dummy import DummyClassifier
+from sklearn.metrics import ConfusionMatrixDisplay, balanced_accuracy_score, classification_report
+
+logger = logging.getLogger(__name__)
+
+
+def evaluate_model(
+ clf: DummyClassifier | CatBoostClassifier,
+ df_eval: pd.DataFrame,
+ feature_cols: list[str],
+ output_path: Path | None = None,
+):
+ logger.info("Evaluating model.")
+ x_eval = np.vstack(df_eval[feature_cols].values)
+ y_pred = clf.predict(x_eval)
+
+ df_eval["y_pred"] = y_pred
+
+ df_eval.loc[df_eval.lang_matches_recertification, ["y_pred"]] = "PREVIOUS_VERSION"
+ df_eval.loc[
+ (df_eval.lang_token_set_ratio == 100)
+ & (df_eval.lang_len_difference < 5)
+ & (df_eval.y_pred != "PREVIOUS_VERSION"),
+ ["y_pred"],
+ ] = "PREVIOUS_VERSION"
+
+ print(classification_report(df_eval.label.values, df_eval.y_pred.values))
+ print(f"Balanced accuracy score: {balanced_accuracy_score(df_eval.label.values, df_eval.y_pred.values)}")
+
+ fig = ConfusionMatrixDisplay.from_predictions(
+ df_eval.label.values,
+ df_eval.y_pred.values,
+ xticks_rotation=90,
+ )
+
+ if output_path:
+ report_dict = classification_report(df_eval.label.values, df_eval.y_pred.values, output_dict=True)
+ report_df = pd.DataFrame(report_dict).transpose()
+ report_df.to_csv(output_path / "classification_report.csv")
+ fig.figure_.savefig(output_path / "confusion_matrix.png")
+ with Path(output_path / "balanced_accuracy_score.txt").open("w") as handle:
+ handle.write(str(balanced_accuracy_score(df_eval.label.values, df_eval.y_pred.values)))
+
+ if isinstance(clf, CatBoostClassifier):
+ feature_importance = clf.get_feature_importance()
+ sorted_idx = np.argsort(feature_importance)
+ features = np.array(feature_cols)[sorted_idx]
+
+ fig_feature_importance = plt.figure(figsize=(10, 12))
+ plt.barh(features, feature_importance[sorted_idx], align="center")
+ plt.xlabel("Feature Importance")
+ plt.ylabel("Feature")
+ plt.title("Feature Importance in Gradient boosted trees classifier")
+ plt.tight_layout()
+ plt.show()
+
+ if output_path:
+ fig_feature_importance.savefig(output_path / "feature_importance.png")
+
+
+def display_dim_red_scatter(df: pd.DataFrame, dim_red: Literal["umap", "pca"]) -> None:
+ df_exploded = df.explode(["segments", dim_red]).reset_index()
+
+ x_col = dim_red + "_x"
+ y_col = dim_red + "_y"
+
+ df_exploded[x_col] = df_exploded[dim_red].map(lambda x: x[0])
+ df_exploded[y_col] = df_exploded[dim_red].map(lambda x: x[1])
+ df_exploded["wrapped_segment"] = df_exploded.segments.str.wrap(60).map(lambda x: x.replace("\n", "<br>"))
+
+ fig = px.scatter(
+ df_exploded,
+ x=x_col,
+ y=y_col,
+ color="label",
+ hover_data=["dgst", "canonical_reference_keyword", "wrapped_segment"],
+ width=1500,
+ height=1000,
+ title=f"{dim_red.upper()} projection of segment embeddings.",
+ )
+ fig.show()
diff --git a/src/sec_certs/model/references_nlp/feature_extraction.py b/src/sec_certs/model/references_nlp/feature_extraction.py
new file mode 100644
index 00000000..eed90959
--- /dev/null
+++ b/src/sec_certs/model/references_nlp/feature_extraction.py
@@ -0,0 +1,542 @@
+import itertools
+import logging
+import re
+from collections import Counter
+from pathlib import Path
+from typing import Literal
+
+import numpy as np
+import pandas as pd
+import spacy
+import umap
+import umap.plot
+from rapidfuzz import fuzz
+from scipy.spatial import ConvexHull, QhullError, distance_matrix
+from scipy.stats import kurtosis, skew
+from sklearn.decomposition import PCA
+from sklearn.feature_extraction.text import TfidfVectorizer
+from sklearn.preprocessing import LabelEncoder, StandardScaler
+
+from sec_certs.constants import RANDOM_STATE, REF_ANNOTATION_MODES, REF_EMBEDDING_METHOD
+from sec_certs.dataset import CCDataset
+from sec_certs.model.references_nlp.annotator import ReferenceAnnotator
+from sec_certs.model.references_nlp.annotator_trainer import ReferenceAnnotatorTrainer
+from sec_certs.model.references_nlp.segment_extractor import ReferenceSegmentExtractor
+from sec_certs.utils.nlp import prec_recall_metric
+
+logger = logging.getLogger(__name__)
+
+nlp = spacy.load("en_core_web_sm")
+
+
+def strip_all(text: str, to_strip) -> str:
+ if pd.isna(to_strip):
+ return text
+ for i in to_strip:
+ text = text.replace(i, "")
+ return text
+
+
+def matches_recertification(segments: list[str]) -> bool:
+ regex_a = r"This is a re-?\s?certification based on (the\s){0,1}REFERENCED_CERTIFICATE_ID"
+ regex_b = r"Re-?\s?Zertifizierung basierend auf (the\s){0,1}REFERENCED_CERTIFICATE_ID"
+ return any(
+ re.search(regex_a, segment, re.IGNORECASE) or re.search(regex_b, segment, re.IGNORECASE) for segment in segments
+ )
+
+
+def compute_ngram_overlap_spacy(string1, string2, n):
+ doc1 = nlp(string1)
+ doc2 = nlp(string2)
+
+ ngrams1 = [" ".join([token.text for token in doc1[i : i + n]]) for i in range(len(doc1) - n + 1)]
+ ngrams2 = [" ".join([token.text for token in doc2[i : i + n]]) for i in range(len(doc2) - n + 1)]
+
+ overlap = sum((Counter(ngrams1) & Counter(ngrams2)).values())
+ return overlap
+
+
+def compute_character_ngram_overlap(str1, str2, n):
+ ngrams1 = [str1[i : i + n] for i in range(len(str1) - n + 1)]
+ ngrams2 = [str2[i : i + n] for i in range(len(str2) - n + 1)]
+ overlap = sum((Counter(ngrams1) & Counter(ngrams2)).values())
+ return overlap
+
+
+def compute_common_length(str1, str2, prefix=True):
+ length = 0
+ min_length = min(len(str1), len(str2))
+ if prefix:
+ for i in range(min_length):
+ if str1[i] == str2[i]:
+ length += 1
+ else:
+ break
+ else:
+ for i in range(1, min_length + 1):
+ if str1[-i] == str2[-i]:
+ length += 1
+ else:
+ break
+ return length
+
+
+def compute_numeric_token_overlap(str1, str2):
+ doc1 = nlp(str1)
+ doc2 = nlp(str2)
+
+ tokens1 = [token.text for token in doc1 if token.like_num]
+ tokens2 = [token.text for token in doc2 if token.like_num]
+
+ overlap = sum((Counter(tokens1) & Counter(tokens2)).values())
+ return overlap
+
+
+def get_lang_features(base_name: str, referenced_name: str) -> tuple:
+ common_numeric_words = compute_numeric_token_overlap(base_name, referenced_name)
+ common_words = compute_ngram_overlap_spacy(base_name, referenced_name, 1)
+ bigram_overlap = compute_ngram_overlap_spacy(base_name, referenced_name, 2)
+ trigram_overlap = compute_ngram_overlap_spacy(base_name, referenced_name, 3)
+ common_prefix_len = compute_common_length(base_name, referenced_name, True)
+ common_suffix_len = compute_common_length(base_name, referenced_name, False)
+ character_bigram_overlap = compute_character_ngram_overlap(base_name, referenced_name, 2)
+ character_trigram_overlap = compute_character_ngram_overlap(base_name, referenced_name, 3)
+ base_len = len(base_name)
+ referenced_len = len(referenced_name)
+ len_difference = abs(base_len - referenced_len)
+
+ return (
+ common_numeric_words,
+ common_words,
+ bigram_overlap,
+ trigram_overlap,
+ common_prefix_len,
+ common_suffix_len,
+ character_bigram_overlap,
+ character_trigram_overlap,
+ base_len,
+ referenced_len,
+ len_difference,
+ )
+
+
+def extract_segments(
+ cc_dset: CCDataset, mode: REF_ANNOTATION_MODES, n_sents_before: int = 2, n_sents_after: int = 1
+) -> pd.DataFrame:
+ logger.info("Extracting segments.")
+ df = ReferenceSegmentExtractor(n_sents_before, n_sents_after)(list(cc_dset.certs.values()))
+ if mode == "training":
+ return df.loc[(df.label.notnull()) & ((df.split == "train") | (df.split == "valid"))]
+ elif mode == "evaluation":
+ return df.loc[df.label.notnull()]
+ elif mode == "production":
+ return df
+ else:
+ raise ValueError(f"Unknown mode {mode}")
+
+
+def _build_transformer_embeddings(
+ segments: pd.DataFrame, mode: REF_ANNOTATION_MODES, model_path: Path | None = None
+) -> pd.DataFrame:
+ should_save_model = model_path is not None
+ annotator = None
+ logger.info("Building transformer embeddings.")
+ if model_path:
+ try:
+ annotator = ReferenceAnnotator.from_pretrained(model_path)
+ should_save_model = False
+ except Exception:
+ print(f"Failed to load ReferenceAnnotator from {model_path}.")
+ should_save_model = True
+
+ if not annotator:
+ print("Training ReferenceAnnotator from scratch.")
+ trainer = ReferenceAnnotatorTrainer.from_df(
+ segments,
+ prec_recall_metric,
+ mode=mode,
+ n_iterations=34,
+ n_epochs=1,
+ learning_rate=0.01,
+ batch_size=16,
+ segmenter_metric="f1",
+ ensemble_soft_voting_power=2,
+ show_progress_bar=False,
+ )
+ trainer.train()
+ annotator = trainer.clf
+ assert annotator is not None
+
+ if should_save_model and model_path:
+ annotator.save_pretrained(model_path)
+
+ return (
+ segments.copy().assign(embeddings=lambda df_: df_.segments.map(annotator._model.model_body.encode)),
+ annotator,
+ )
+
+
+def _build_tf_idf_embeddings(segments: pd.DataFrame, mode: REF_ANNOTATION_MODES) -> pd.DataFrame:
+ def choose_values_to_fit(df_: pd.DataFrame) -> list[str]:
+ if mode == "training":
+ return df_.loc[df_.split == "train"].copy().explode("segments").segments.values
+ elif mode == "evaluation":
+ return df_.loc[df_.split != "test"].copy().explode("segments").segments.values
+ elif mode == "production":
+ return df_.copy().explode("segments").segments.values
+ else:
+ raise ValueError(f"Unknown mode {mode}")
+
+ logger.info("Building TF-IDF embeddings.")
+ tf_idf = TfidfVectorizer()
+ tf_idf = tf_idf.fit(choose_values_to_fit(segments))
+
+ return segments.copy().assign(
+ embeddings=lambda df_: df_.segments.map(lambda x: tf_idf.transform(x).toarray().tolist())
+ )
+
+
+def build_embeddings(
+ segments: pd.DataFrame, mode: REF_ANNOTATION_MODES, method: REF_EMBEDDING_METHOD, model_path: Path | None = None
+) -> pd.DataFrame:
+ return (
+ _build_transformer_embeddings(segments, mode, model_path)
+ if method == "transformer"
+ else _build_tf_idf_embeddings(segments, mode)
+ )
+
+
+def extract_language_features(df: pd.DataFrame, cc_dset: CCDataset) -> pd.DataFrame:
+ logger.info("Extracting language features.")
+ certs = list(cc_dset.certs.values())
+ dgst_to_cert_name = {x.dgst: x.name for x in certs}
+ cert_id_to_cert_name = {x.heuristics.cert_id: x.name for x in certs}
+ dgst_to_extracted_versions = {x.dgst: x.heuristics.extracted_versions for x in certs}
+ cert_id_to_extracted_versions = {x.heuristics.cert_id: x.heuristics.extracted_versions for x in certs}
+
+ df_lang = (
+ df.copy()
+ .assign(
+ cert_name=lambda df_: df_.dgst.map(dgst_to_cert_name),
+ referenced_cert_name=lambda df_: df_.canonical_reference_keyword.map(cert_id_to_cert_name),
+ cert_versions=lambda df_: df_.dgst.map(dgst_to_extracted_versions),
+ referenced_cert_versions=lambda df_: df_.canonical_reference_keyword.map(cert_id_to_extracted_versions),
+ cert_name_stripped_version=lambda df_: df_.apply(
+ lambda x: strip_all(x["cert_name"], x["cert_versions"]), axis=1
+ ),
+ referenced_cert_name_stripped_version=lambda df_: df_.apply(
+ lambda x: strip_all(x["referenced_cert_name"], x["referenced_cert_versions"]), axis=1
+ ),
+ lang_token_set_ratio=lambda df_: df_.apply(
+ lambda x: fuzz.token_set_ratio(
+ x["cert_name_stripped_version"], x["referenced_cert_name_stripped_version"]
+ ),
+ axis=1,
+ ),
+ lang_partial_ratio=lambda df_: df_.apply(
+ lambda x: fuzz.partial_ratio(
+ x["cert_name_stripped_version"], x["referenced_cert_name_stripped_version"]
+ ),
+ axis=1,
+ ),
+ lang_token_sort_ratio=lambda df_: df_.apply(
+ lambda x: fuzz.token_sort_ratio(
+ x["cert_name_stripped_version"], x["referenced_cert_name_stripped_version"]
+ ),
+ axis=1,
+ ),
+ lang_n_segments=lambda df_: df_.segments.map(lambda x: len(x) if x else 0),
+ lang_matches_recertification=lambda df_: df_.segments.map(matches_recertification),
+ )
+ .assign(
+ lang_n_extracted_versions=lambda df_: df_.cert_versions.map(lambda x: len(x) if x else 0),
+ lang_n_intersection_versions=lambda df_: df_.apply(
+ lambda x: len(set(x["cert_versions"]).intersection(set(x["referenced_cert_versions"]))), axis=1
+ ),
+ )
+ )
+
+ df_lang_other_features = df_lang.apply(
+ lambda row: get_lang_features(row["cert_name"], row["referenced_cert_name"]), axis=1
+ ).apply(pd.Series)
+ lang_features = [
+ "common_numeric_words",
+ "common_words",
+ "bigram_overlap",
+ "trigram_overlap",
+ "common_prefix_len",
+ "common_suffix_len",
+ "character_bigram_overlap",
+ "character_trigram_overlap",
+ "base_len",
+ "referenced_len",
+ "len_difference",
+ ]
+ df_lang_other_features.columns = ["lang_" + x for x in lang_features]
+
+ df_lang = pd.concat([df_lang, df_lang_other_features], axis=1).assign(
+ lang_should_not_be_component=lambda df_: df_.apply(
+ lambda x: x.lang_len_difference < 5 and x.lang_token_set_ratio == 100, axis=1
+ ),
+ )
+ for col in df_lang.columns:
+ if col.startswith("pred_"):
+ df_lang[col] = df_lang[col] / df_lang.lang_n_segments
+
+ return df_lang
+
+
+def perform_dimensionality_reduction(
+ df: pd.DataFrame,
+ mode: REF_ANNOTATION_MODES,
+ umap_n_neighbors: int = 5,
+ umap_min_dist: float = 0.1,
+ umap_metric: Literal["cosine", "euclidean", "manhattan"] = "euclidean",
+) -> pd.DataFrame:
+ def choose_values_to_fit(df_: pd.DataFrame):
+ if mode == "training":
+ return df_.loc[df_.split == "train"].copy().embeddings.values
+ elif mode == "evaluation":
+ return df_.loc[df_.split != "test"].copy().embeddings.values
+ elif mode == "production":
+ return df_.copy().embeddings.values
+ else:
+ raise ValueError(f"Unknown mode {mode}")
+
+ def choose_labels_to_fit(df_: pd.DataFrame):
+ if mode == "training":
+ return df_.loc[df_.split == "train"].copy().label.values
+ elif mode == "evaluation":
+ return df_.loc[df_.split != "test"].copy().label.values
+ elif mode == "production":
+ return df_.copy().label.values
+ else:
+ raise ValueError(f"Unknown mode {mode}")
+
+ logger.info("Performing dimensionality reduction.")
+ df_exploded = df.copy().explode(["segments", "embeddings"]).reset_index(drop=True)
+ label_encoder = LabelEncoder()
+
+ embeddings_to_fit = np.vstack(choose_values_to_fit(df_exploded))
+ labels_to_fit = label_encoder.fit_transform(choose_labels_to_fit(df_exploded))
+
+ scaler = StandardScaler()
+ embeddings_to_fit_scaled = scaler.fit_transform(embeddings_to_fit)
+
+ # parallel UMAP not available with random state
+ umapper = umap.UMAP(
+ n_neighbors=umap_n_neighbors, min_dist=umap_min_dist, metric=umap_metric, random_state=RANDOM_STATE, n_jobs=1
+ ).fit(embeddings_to_fit, y=labels_to_fit)
+ pca_mapper = PCA(n_components=2, random_state=RANDOM_STATE).fit(embeddings_to_fit_scaled, y=labels_to_fit)
+
+ all_embeddings = np.vstack(df.embeddings.values)
+ all_embeddings_scaled = scaler.transform(all_embeddings)
+
+ df_exploded["umap"] = umapper.transform(all_embeddings).tolist()
+ df_exploded["pca"] = pca_mapper.transform(all_embeddings_scaled).tolist()
+
+ return (
+ df_exploded.groupby(["dgst", "canonical_reference_keyword"])
+ .agg(
+ {
+ "segments": lambda x: x.tolist(),
+ "actual_reference_keywords": "first",
+ "label": "first",
+ "split": "first",
+ "embeddings": lambda x: x.tolist(),
+ "umap": lambda x: x.tolist(),
+ "pca": lambda x: x.tolist(),
+ }
+ )
+ .reset_index()
+ )
+
+
+def extract_prediction_features(df: pd.DataFrame, model) -> pd.DataFrame:
+ def get_setfit_prediction_numbers(val):
+ counter = Counter(val.tolist())
+ return [counter[x] for x in range(len(all_labels))]
+
+ logger.info("Extracting prediction features.")
+ df["annotator_predictions"] = df.segments.map(lambda x: model.predict(x))
+ all_labels = set(itertools.chain.from_iterable(x.tolist() for x in df.annotator_predictions.values))
+
+ df_features_pred = df.annotator_predictions.apply(get_setfit_prediction_numbers).apply(pd.Series)
+ feature_names = [f"pred_{x}" for x in range(len(all_labels))]
+ df_features_pred.columns = feature_names
+ return pd.concat([df, df_features_pred], axis=1)
+
+
+def extract_geometrical_features(df: pd.DataFrame) -> pd.DataFrame:
+ def extract_features(points):
+ # Convert list of points to a numpy array
+ points = np.array(points)
+ xs = points[:, 0]
+ ys = points[:, 1]
+
+ # Basic Descriptive Statistics
+ mean_x, mean_y = np.mean(xs), np.mean(ys)
+ var_x, var_y = np.var(xs), np.var(ys)
+ std_x, std_y = np.std(xs), np.std(ys)
+ if len(points) > 1:
+ skew_x, skew_y = skew(xs), skew(ys)
+ kurt_x, kurt_y = kurtosis(xs), kurtosis(ys)
+ else:
+ skew_x, skew_y = 0, 0
+ kurt_x, kurt_y = 0, 0
+
+ # Spatial Spread
+ range_x, range_y = np.ptp(xs), np.ptp(ys)
+ cov_xy = np.cov(xs, ys)[0, 1] if len(points) > 1 else 0
+ median_x, median_y = np.median(xs), np.median(ys)
+
+ # Distance-based Features
+ centroid = [mean_x, mean_y]
+ distances_to_centroid = np.linalg.norm(points - centroid, axis=1) if len(points) > 1 else [0]
+ mean_distance = np.mean(distances_to_centroid)
+ max_distance = np.max(distances_to_centroid)
+ min_distance = np.min(distances_to_centroid)
+ std_distance = np.std(distances_to_centroid)
+ max_min_distance = max_distance - min_distance
+
+ sorted_points = points[np.argsort(distances_to_centroid)]
+ total_distance = np.sum(np.linalg.norm(sorted_points[1:] - sorted_points[:-1], axis=1))
+
+ # Geometric Features
+ hull_area, hull_perimeter = (0, 0)
+ if len(points) > 2: # ConvexHull needs at least 3 points
+ try:
+ hull = ConvexHull(points)
+ hull_area = hull.volume
+ hull_perimeter = hull.area
+ except QhullError:
+ pass
+
+ pairwise_distances = distance_matrix(points, points) if len(points) > 1 else np.array([[0]])
+ mean_pairwise_distance = np.mean(pairwise_distances)
+ max_pairwise_distance = np.max(pairwise_distances)
+
+ if len(points) > 1:
+ min_coords = np.min(points, axis=0)
+ max_coords = np.max(points, axis=0)
+ bounding_box_width = max_coords[0] - min_coords[0]
+ bounding_box_height = max_coords[1] - min_coords[1]
+ bounding_box_area = bounding_box_width * bounding_box_height
+
+ aspect_ratio = bounding_box_width / bounding_box_height if bounding_box_height != 0 else 1
+ point_density = len(points) / bounding_box_area
+ else:
+ aspect_ratio = 0
+ point_density = 0
+
+ # Gather all features into a list
+ features = [
+ mean_x,
+ mean_y,
+ var_x,
+ var_y,
+ std_x,
+ std_y,
+ skew_x,
+ skew_y,
+ kurt_x,
+ kurt_y,
+ range_x,
+ range_y,
+ cov_xy,
+ median_x,
+ median_y,
+ mean_distance,
+ max_distance,
+ min_distance,
+ max_min_distance,
+ std_distance,
+ total_distance,
+ hull_area,
+ hull_perimeter,
+ mean_pairwise_distance,
+ max_pairwise_distance,
+ aspect_ratio,
+ point_density,
+ ]
+
+ return features
+
+ feature_names = [
+ "mean_x",
+ "mean_y",
+ "var_x",
+ "var_y",
+ "std_x",
+ "std_y",
+ "skew_x",
+ "skew_y",
+ "kurt_x",
+ "kurt_y",
+ "range_x",
+ "range_y",
+ "cov_xy",
+ "median_x",
+ "median_y",
+ "mean_distance_to_centroid",
+ "max_distance_to_centroid",
+ "min_distance_to_centroid",
+ "max_min_distance_to_centroid",
+ "std_distance_to_centroid",
+ "total_distances_to_centroid",
+ "hull_area",
+ "hull_perimeter",
+ "mean_pairwise_distance",
+ "max_pairwise_distance",
+ "aspect_ratio",
+ "point_density",
+ ]
+
+ logger.info("Extracting geometrical features.")
+ df_features_pca = df.pca.apply(extract_features).apply(pd.Series)
+ feature_names_pca = ["pca_" + x for x in feature_names]
+ df_features_pca.columns = feature_names_pca
+
+ df_features_umap = df.umap.apply(extract_features).apply(pd.Series)
+ feature_names_umap = ["umap_" + x for x in feature_names]
+ df_features_umap.columns = feature_names_umap
+
+ return pd.concat([df, df_features_pca, df_features_umap], axis=1)
+
+
+def get_data_for_clf(
+ df: pd.DataFrame,
+ mode: REF_ANNOTATION_MODES,
+ use_pca: bool = True,
+ use_umap: bool = True,
+ use_lang: bool = True,
+ use_pred: bool = True,
+) -> tuple[np.ndarray, np.ndarray, pd.DataFrame | None, list[str]]:
+ feature_columns = []
+ if not use_pca and not use_umap and not use_lang and not use_pred:
+ raise ValueError("At least one of PCA, UMAP or language features must be used.")
+ if use_pca:
+ feature_columns.extend([x for x in df.columns if x.startswith("pca_")])
+ if use_umap:
+ feature_columns.extend([x for x in df.columns if x.startswith("umap_")])
+ if use_lang:
+ feature_columns.extend([x for x in df.columns if x.startswith("lang_")])
+ if use_pred:
+ feature_columns.extend([x for x in df.columns if x.startswith("pred_")])
+
+ if mode == "training":
+ train_df = df.loc[df.split == "train"].copy()
+ eval_df = df.loc[df.split == "valid"].copy()
+ elif mode == "evaluation":
+ train_df = df.loc[df.split != "test"].copy()
+ eval_df = df.loc[df.split == "test"].copy()
+ elif mode == "production":
+ train_df = df.copy()
+ eval_df = df.copy()
+ elif mode == "cross-validation":
+ train_df = df.loc[df.split != "test"].copy()
+ eval_df = None
+ else:
+ raise ValueError(f"Unknown mode {mode}")
+
+ return np.vstack(train_df[feature_columns].values), train_df.label.values, eval_df, feature_columns
diff --git a/src/sec_certs/model/references/segment_extractor.py b/src/sec_certs/model/references_nlp/segment_extractor.py
index bf6096bc..ce144b93 100644
--- a/src/sec_certs/model/references/segment_extractor.py
+++ b/src/sec_certs/model/references_nlp/segment_extractor.py
@@ -13,7 +13,6 @@ from typing import Any, Iterable, Literal
import numpy as np
import pandas as pd
import spacy
-from rapidfuzz import fuzz
from sec_certs.sample.cc import CCCertificate
from sec_certs.sample.cc_certificate_id import CertificateId
@@ -81,14 +80,6 @@ def preprocess_data_source(record: ReferenceRecord) -> ReferenceRecord:
return record
-def strip_all(text: str, to_strip) -> str:
- if pd.isna(to_strip):
- return text
- for i in to_strip:
- text = text.replace(i, "")
- return text
-
-
def find_bracket_pattern(sentences: set[str], actual_reference_keywords: frozenset[str]):
patterns = [r"(\[.+?\])(?=.*" + x + r")" for x in actual_reference_keywords]
res: list[tuple[str, str]] = []
@@ -166,8 +157,9 @@ class ReferenceSegmentExtractor:
Should be only called with ReferenceSegmentExtractor()(list_of_certificates)
"""
- def __init__(self):
- pass
+ def __init__(self, n_sents_before: int = 1, n_sents_after: int = 0):
+ self.n_sents_before = n_sents_before
+ self.n_sents_after = n_sents_after
def __call__(self, certs: Iterable[CCCertificate]) -> pd.DataFrame:
return self._prepare_df_from_cc_dset(certs)
@@ -235,10 +227,12 @@ class ReferenceSegmentExtractor:
progress_bar=True,
progress_bar_desc="Preprocessing data",
)
+ records_with_args = [(x, self.n_sents_before, self.n_sents_after) for x in records]
results = parallel_processing.process_parallel(
fill_reference_segments,
- records,
+ records_with_args,
+ unpack=True,
use_threading=False,
progress_bar=True,
progress_bar_desc="Recovering reference segments",
@@ -322,13 +316,6 @@ class ReferenceSegmentExtractor:
"""
annotations_dict = ReferenceSegmentExtractor._get_annotations_dict()
split_dct = ReferenceSegmentExtractor._get_split_dict()
-
- # Retrieve some columns previously lost
- dgst_to_cert_name = {x.dgst: x.name for x in certs}
- cert_id_to_cert_name = {x.heuristics.cert_id: x.name for x in certs}
- dgst_to_extracted_versions = {x.dgst: x.heuristics.extracted_versions for x in certs}
- cert_id_to_extracted_versions = {x.heuristics.cert_id: x.heuristics.extracted_versions for x in certs}
-
logger.info(f"Deleting {df.loc[df.segments.isnull()].shape[0]} rows with no segments.")
df_new = df.copy()
@@ -350,34 +337,11 @@ class ReferenceSegmentExtractor:
)
.agg({"segments": list, "actual_reference_keywords": unique_elements})
.assign(
- split=lambda df_: df_.dgst.map(split_dct),
+ actual_reference_keywords=lambda df_: df_.actual_reference_keywords.map(list),
label=lambda df_: [
annotations_dict.get(x) for x in zip(df_["dgst"], df_["canonical_reference_keyword"])
],
- actual_reference_keywords=lambda df_: df_.actual_reference_keywords.map(list),
- cert_name=lambda df_: df_.dgst.map(dgst_to_cert_name),
- referenced_cert_name=lambda df_: df_.canonical_reference_keyword.map(cert_id_to_cert_name),
- cert_versions=lambda df_: df_.dgst.map(dgst_to_extracted_versions),
- referenced_cert_versions=lambda df_: df_.canonical_reference_keyword.map(cert_id_to_extracted_versions),
- cert_name_stripped_version=lambda df_: df_.apply(
- lambda x: strip_all(x["cert_name"], x["cert_versions"]), axis=1
- ),
- referenced_cert_name_stripped_version=lambda df_: df_.apply(
- lambda x: strip_all(x["referenced_cert_name"], x["referenced_cert_versions"]), axis=1
- ),
- name_similarity=lambda df_: df_.apply(
- lambda x: fuzz.token_set_ratio(
- x["cert_name_stripped_version"], x["referenced_cert_name_stripped_version"]
- ),
- axis=1,
- ),
- name_len_diff=lambda df_: df_.apply(
- lambda x: np.nan
- if pd.isnull(x["cert_name_stripped_version"])
- or pd.isnull(x["referenced_cert_name_stripped_version"])
- else abs(len(x["cert_name_stripped_version"]) - len(x["referenced_cert_name_stripped_version"])),
- axis=1,
- ),
+ split=lambda df_: df_.dgst.map(split_dct),
)
.assign(
label=lambda df_: df_.label.map(lambda x: x if x is not None else np.nan),
diff --git a/src/sec_certs/model/references_nlp/training.py b/src/sec_certs/model/references_nlp/training.py
new file mode 100644
index 00000000..4e0f82ce
--- /dev/null
+++ b/src/sec_certs/model/references_nlp/training.py
@@ -0,0 +1,68 @@
+import logging
+import os
+
+import numpy as np
+import pandas as pd
+from catboost import CatBoostClassifier, Pool
+from sklearn.dummy import DummyClassifier
+from sklearn.metrics import balanced_accuracy_score
+from sklearn.model_selection import KFold
+
+from sec_certs.constants import RANDOM_STATE, REF_ANNOTATION_MODES
+from sec_certs.model.references_nlp.feature_extraction import get_data_for_clf
+
+logger = logging.getLogger(__name__)
+
+
+def _train_model(x_train, y_train, x_eval, y_eval, learning_rate: float = 0.03, depth: int = 6, l2_leaf_reg: int = 3):
+ clf = CatBoostClassifier(
+ learning_rate=learning_rate,
+ depth=depth,
+ l2_leaf_reg=l2_leaf_reg,
+ task_type="GPU",
+ devices=os.environ["CUDA_VISIBLE_DEVICES"],
+ random_seed=RANDOM_STATE,
+ )
+
+ train_pool = Pool(x_train, y_train)
+ eval_pool = Pool(x_eval, y_eval)
+ clf.fit(train_pool, eval_set=eval_pool, verbose=False, plot=True, early_stopping_rounds=100, use_best_model=True)
+ return clf
+
+
+def train_model(
+ df: pd.DataFrame,
+ mode: REF_ANNOTATION_MODES,
+ train_baseline: bool = False,
+ use_pca: bool = True,
+ use_umap: bool = True,
+ use_lang: bool = True,
+ use_pred: bool = True,
+ learning_rate: float = 0.03,
+ depth: int = 6,
+ l2_leaf_reg: int = 3,
+) -> tuple[DummyClassifier | CatBoostClassifier, pd.DataFrame, list[str]]:
+ logger.info(f"Training model for mode {mode}")
+ X_train, y_train, eval_df, feature_cols = get_data_for_clf(df, mode, use_pca, use_umap, use_lang, use_pred)
+ if train_baseline:
+ clf = DummyClassifier(random_state=RANDOM_STATE)
+ clf.fit(X_train, y_train)
+ else:
+ assert eval_df is not None
+ clf = _train_model(X_train, y_train, eval_df[feature_cols], eval_df.label, learning_rate, depth, l2_leaf_reg)
+
+ return clf, eval_df, feature_cols
+
+
+def cross_validate_model(df: pd.DataFrame, learning_rate: float = 0.03, depth: int = 6, l2_leaf_reg: int = 3) -> float:
+ logger.info("Cross-validating model")
+ X_train, y_train, _, _ = get_data_for_clf(df, "cross-validation", True, True, True, True)
+ kf = KFold(n_splits=5, shuffle=True, random_state=RANDOM_STATE)
+ scores = []
+ for train_index, test_index in kf.split(X_train):
+ X_train_, X_test_ = X_train[train_index], X_train[test_index]
+ y_train_, y_test_ = y_train[train_index], y_train[test_index]
+ clf = _train_model(X_train_, y_train_, X_test_, y_test_, learning_rate, depth, l2_leaf_reg)
+ scores.append(balanced_accuracy_score(y_test_, clf.predict(X_test_)))
+
+ return np.mean(scores)