aboutsummaryrefslogtreecommitdiffhomepage
path: root/src/sec_certs/utils
diff options
context:
space:
mode:
authoradamjanovsky2023-03-03 14:55:26 +0100
committeradamjanovsky2023-03-03 14:55:26 +0100
commit81cde7965737006e183a98251f515fac8c0df13d (patch)
treea957711e381902ab06dfe37ddf1d6c159f2b36d1 /src/sec_certs/utils
parenta53f0f71ab0c741d274693aa03e34c74df10d5a2 (diff)
downloadsec-certs-81cde7965737006e183a98251f515fac8c0df13d.tar.gz
sec-certs-81cde7965737006e183a98251f515fac8c0df13d.tar.zst
sec-certs-81cde7965737006e183a98251f515fac8c0df13d.zip
WiP production-level reference annotation
Diffstat (limited to 'src/sec_certs/utils')
-rw-r--r--src/sec_certs/utils/nlp.py24
1 files changed, 24 insertions, 0 deletions
diff --git a/src/sec_certs/utils/nlp.py b/src/sec_certs/utils/nlp.py
index 496b71b8..45ad9639 100644
--- a/src/sec_certs/utils/nlp.py
+++ b/src/sec_certs/utils/nlp.py
@@ -1,4 +1,9 @@
+from __future__ import annotations
+
+from ast import literal_eval
+
import numpy as np
+import pandas as pd
from sklearn.metrics import precision_score, recall_score
@@ -11,3 +16,22 @@ def prec_recall_metric(y_pred, y_true):
def softmax(x):
return np.exp(x - np.max(x)) / np.exp(x - np.max(x)).sum()
+
+
+def eval_strings(series):
+ return [list(literal_eval(x)) for x in series]
+
+
+def filter_short_sentences(sentences, cert_id):
+ return [x for x in sentences if len(x) > len(cert_id) + 20]
+
+
+def prepare_reference_annotations_df(df: pd.DataFrame):
+ df = (
+ df.loc[lambda df_: (df_.label != "SELF") & (df_.label.notnull())]
+ .assign(segments=lambda df_: eval_strings(df_.segments))
+ .drop(columns="lang")
+ )
+ df.segments = df.apply(lambda row: filter_short_sentences(row["segments"], row["referenced_cert_id"]), axis=1)
+ df = df.loc[lambda df_: df_.segments.map(len) > 0]
+ return df