diff options
| author | adamjanovsky | 2023-11-14 10:04:13 +0100 |
|---|---|---|
| committer | adamjanovsky | 2023-11-14 10:04:13 +0100 |
| commit | 80190b01aeda844b9d3ea8684284130c44f1453e (patch) | |
| tree | 6fbcabd9cda272b9a5d64c8e61c7d3b914351f93 /src | |
| parent | 9cdf4801f93243e682b43be0a52956c0f9fad377 (diff) | |
| download | sec-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.py | 7 | ||||
| -rw-r--r-- | src/sec_certs/data/reference_annotations/readme.md | 31 | ||||
| -rw-r--r-- | src/sec_certs/dataset/cc.py | 69 | ||||
| -rw-r--r-- | src/sec_certs/dataset/fips.py | 2 | ||||
| -rw-r--r-- | src/sec_certs/model/__init__.py | 5 | ||||
| -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__.py | 13 | ||||
| -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.py | 90 | ||||
| -rw-r--r-- | src/sec_certs/model/references_nlp/feature_extraction.py | 542 | ||||
| -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.py | 68 |
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) |
