aboutsummaryrefslogtreecommitdiffhomepage
diff options
context:
space:
mode:
authoradamjanovsky2023-12-18 13:03:23 +0100
committeradamjanovsky2023-12-18 13:03:23 +0100
commit1ad29f2a34178550e1ddcdc617899e00866903c5 (patch)
tree353d73122c2f3d3c7fb2f0fcddcf347728a614cc
parentbfffc154e17349751e93bc0d89f5926f64680d5f (diff)
downloadsec-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.ipynb42
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()"
]
}
],