aboutsummaryrefslogtreecommitdiffhomepage
diff options
context:
space:
mode:
authoradamjanovsky2023-03-09 13:45:55 +0100
committeradamjanovsky2023-03-09 13:45:55 +0100
commit0787c7e2fb4472e56af887c8e2168f0aa6490e23 (patch)
treedb9d5f70f8e6389555c495b44f4dc12f6f74f6dc
parent81cde7965737006e183a98251f515fac8c0df13d (diff)
downloadsec-certs-0787c7e2fb4472e56af887c8e2168f0aa6490e23.tar.gz
sec-certs-0787c7e2fb4472e56af887c8e2168f0aa6490e23.tar.zst
sec-certs-0787c7e2fb4472e56af887c8e2168f0aa6490e23.zip
minor improvements
-rw-r--r--src/sec_certs/model/references/annotator.py2
-rw-r--r--src/sec_certs/model/references/annotator_trainer.py4
-rw-r--r--src/sec_certs/model/references/segment_extractor.py23
3 files changed, 14 insertions, 15 deletions
diff --git a/src/sec_certs/model/references/annotator.py b/src/sec_certs/model/references/annotator.py
index 4b02430c..e64a7221 100644
--- a/src/sec_certs/model/references/annotator.py
+++ b/src/sec_certs/model/references/annotator.py
@@ -51,7 +51,7 @@ class ReferenceAnnotator:
self._model.save_pretrained(str(model_dir))
def train(self, train_dataset: pd.DataFrame):
- pass
+ raise NotImplementedError("ReferenceAnnotatorTrainer shall be used for training")
def predict(self, X: list[list[str]]) -> list[str]:
return [self._predict_single(x) for x in X]
diff --git a/src/sec_certs/model/references/annotator_trainer.py b/src/sec_certs/model/references/annotator_trainer.py
index 4ee712fc..77c55754 100644
--- a/src/sec_certs/model/references/annotator_trainer.py
+++ b/src/sec_certs/model/references/annotator_trainer.py
@@ -72,12 +72,12 @@ class ReferenceAnnotatorTrainer:
mode: Literal["training", "production"] = "training",
):
df = prepare_reference_annotations_df(df)
- processing_method = {
+ dataset_generation_method = {
"training": ReferenceAnnotatorTrainer.split_df_for_training,
"production": ReferenceAnnotatorTrainer.split_df_for_production,
}
- train_dataset, eval_dataset = processing_method[mode](df)
+ train_dataset, eval_dataset = dataset_generation_method[mode](df)
return cls(train_dataset, eval_dataset, metric, method)
@staticmethod
diff --git a/src/sec_certs/model/references/segment_extractor.py b/src/sec_certs/model/references/segment_extractor.py
index f2496447..23ff4c1d 100644
--- a/src/sec_certs/model/references/segment_extractor.py
+++ b/src/sec_certs/model/references/segment_extractor.py
@@ -43,8 +43,10 @@ class ReferenceRecord:
data = handle.read()
record.segments = {sent.text for sent in nlp(data).sents if record.referenced_cert_id in sent.text}
+
if not record.segments:
record.segments = None
+
return record
def to_pandas_tuple(self) -> tuple[str, str, str, set[str] | None]:
@@ -68,18 +70,16 @@ class ReferenceSegmentExtractor:
- Loads manually annotated samples
- Combines all of that into single dataframe
"""
- df_targets = self._build_df(
- [x for x in certs if x.heuristics.st_references.directly_referencing and x.state.st_txt_path], "target"
- )
- df_reports = self._build_df(
- [x for x in certs if x.heuristics.report_references.directly_referencing and x.state.report_txt_path],
- "report",
- )
- df = pd.concat([df_targets, df_reports])
- return self._process_df(df)
+
+ target_certs = [x for x in certs if x.heuristics.st_references.directly_referencing and x.state.st_txt_path]
+ report_certs = [
+ x for x in certs if x.heuristics.report_references.directly_referencing and x.state.report_txt_path
+ ]
+ df_targets = self._build_df(target_certs, "target")
+ df_reports = self._build_df(report_certs, "report")
+ return self._process_df(pd.concat([df_targets, df_reports]))
def _build_df(self, certs: list[CCCertificate], source: Literal["target", "report"]) -> pd.DataFrame:
- """ """
attribute_mapping = {"target": "st_references", "report": "report_references"}
records = [
ReferenceRecord(x, y, source)
@@ -87,7 +87,6 @@ class ReferenceSegmentExtractor:
for y in getattr(x.heuristics, attribute_mapping[source]).directly_referencing
]
- # results = [ReferenceRecord.fill_reference_segments(x) for x in tqdm.tqdm(records)]
results = parallel_processing.process_parallel(
ReferenceRecord.fill_reference_segments,
records,
@@ -137,7 +136,7 @@ class ReferenceSegmentExtractor:
load_single_df(annotations_directory / "valid.csv", "valid"),
load_single_df(annotations_directory / "test.csv", "test"),
]
- )[["dgst", "referenced_cert_id", "source", "label", "comment"]]
+ )
return (
df_annot[["dgst", "referenced_cert_id", "label"]].set_index(["dgst", "referenced_cert_id"]).label.to_dict()