diff options
| author | adamjanovsky | 2023-03-09 13:45:55 +0100 |
|---|---|---|
| committer | adamjanovsky | 2023-03-09 13:45:55 +0100 |
| commit | 0787c7e2fb4472e56af887c8e2168f0aa6490e23 (patch) | |
| tree | db9d5f70f8e6389555c495b44f4dc12f6f74f6dc | |
| parent | 81cde7965737006e183a98251f515fac8c0df13d (diff) | |
| download | sec-certs-0787c7e2fb4472e56af887c8e2168f0aa6490e23.tar.gz sec-certs-0787c7e2fb4472e56af887c8e2168f0aa6490e23.tar.zst sec-certs-0787c7e2fb4472e56af887c8e2168f0aa6490e23.zip | |
minor improvements
| -rw-r--r-- | src/sec_certs/model/references/annotator.py | 2 | ||||
| -rw-r--r-- | src/sec_certs/model/references/annotator_trainer.py | 4 | ||||
| -rw-r--r-- | src/sec_certs/model/references/segment_extractor.py | 23 |
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() |
