diff options
| author | adamjanovsky | 2023-12-18 13:03:23 +0100 |
|---|---|---|
| committer | adamjanovsky | 2023-12-18 13:03:23 +0100 |
| commit | 1ad29f2a34178550e1ddcdc617899e00866903c5 (patch) | |
| tree | 353d73122c2f3d3c7fb2f0fcddcf347728a614cc | |
| parent | bfffc154e17349751e93bc0d89f5926f64680d5f (diff) | |
| download | sec-certs-1ad29f2a34178550e1ddcdc617899e00866903c5.tar.gz sec-certs-1ad29f2a34178550e1ddcdc617899e00866903c5.tar.zst sec-certs-1ad29f2a34178550e1ddcdc617899e00866903c5.zip | |
compare ref. annotation misclassifications per scheme
| -rw-r--r-- | notebooks/cc/reference_annotations/prediction.ipynb | 42 |
1 files changed, 27 insertions, 15 deletions
diff --git a/notebooks/cc/reference_annotations/prediction.ipynb b/notebooks/cc/reference_annotations/prediction.ipynb index 4eec5da8..940a6b08 100644 --- a/notebooks/cc/reference_annotations/prediction.ipynb +++ b/notebooks/cc/reference_annotations/prediction.ipynb @@ -42,7 +42,7 @@ "logging.getLogger(\"sentence_transformers\").setLevel(logging.CRITICAL)\n", "file_handler = logging.StreamHandler(sys.stderr)\n", "file_handler.setFormatter(logging.Formatter(\"%(asctime)s - %(name)s - %(levelname)s - %(message)s\"))\n", - "logging.basicConfig(level=logging.INFO, handlers=[file_handler])\n" + "logging.basicConfig(level=logging.INFO, handlers=[file_handler])" ] }, { @@ -51,7 +51,7 @@ "metadata": {}, "outputs": [], "source": [ - "mode = \"production\"\n", + "mode = \"evaluation\"\n", "cc_dset = CCDataset.from_json(DATASET_PATH)\n", "\n", "# df = extract_segments(cc_dset, mode=mode)\n", @@ -151,7 +151,7 @@ "df[\"reference_label\"] = df.label.fillna(df.y_pred)\n", "df[[\"dgst\", \"canonical_reference_keyword\", \"reference_label\"]].to_csv(\n", " \"/var/tmp/xjanovsk/certs/sec-certs/dataset/reference_prediction/predictions.csv\"\n", - ")\n" + ")" ] }, { @@ -192,7 +192,7 @@ " y_valid,\n", " features,\n", " output_path=Path(\"/var/tmp/xjanovsk/certs/sec-certs/dataset/cc_ref_annotator_evaluation/embeddings\"),\n", - ")\n" + ")" ] }, { @@ -208,8 +208,10 @@ "metadata": {}, "outputs": [], "source": [ - "misclassified_instances = df.loc[df.y_pred != df.label]\n", - "misclassified_instances = misclassified_instances[\n", + "scheme_mapping = {x.dgst: x.scheme for x in cc_dset}\n", + "all_classified_instances = df.loc[(df.label.notnull())].assign(scheme=lambda df_: df_.dgst.map(scheme_mapping))\n", + "misclassified_instances = df.loc[\n", + " (df.y_pred != df.label) & (df.label.notnull()),\n", " [\n", " \"dgst\",\n", " \"canonical_reference_keyword\",\n", @@ -223,20 +225,30 @@ " \"referenced_cert_versions\",\n", " \"lang_partial_ratio\",\n", " \"lang_token_sort_ratio\",\n", - " ]\n", - "]\n", - "misclassified_instances[\"report_link\"] = misclassified_instances.dgst.map(\n", - " lambda x: f\"https://seccerts.org/cc/{x}/report.pdf\"\n", - ")\n", - "misclassified_instances[\"st_link\"] = misclassified_instances.dgst.map(\n", - " lambda x: f\"https://seccerts.org/cc/{x}/target.pdf\"\n", + " ],\n", + "].assign(\n", + " report_link=lambda df_: df_.dgst.map(lambda x: f\"https://seccerts.org/cc/{x}/report.pdf\"),\n", + " st_link=lambda df_: df_.dgst.map(lambda x: f\"https://seccerts.org/cc/{x}/target.pdf\"),\n", + " scheme=lambda df_: df_.dgst.map(scheme_mapping),\n", ")\n", + "\n", "# Then replace all \\\\/ with / in the corresponding json, as the pandas to_json method escapes the slashes.\n", "misclassified_instances.to_json(\n", - " \"/var/tmp/xjanovsk/certs/sec-certs/dataset/misclassified_references_validation_set.json\",\n", + " REPO_ROOT / \"dataset/misclassified_references_validation_set.json\",\n", " orient=\"records\",\n", " indent=4,\n", - ")\n" + ")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# Display proportion of misclassifications per scheme. Only DE and FR have sufficient support to make any conclusions.\n", + "# FR 4times more likely to be misclassified than DE\n", + "misclassified_instances.scheme.value_counts() * 100 / all_classified_instances.scheme.value_counts()" ] } ], |
