diff --git a/spacy_pii_pipeline.ipynb b/spacy_pii_pipeline.ipynb new file mode 100644 index 0000000..8a99951 --- /dev/null +++ b/spacy_pii_pipeline.ipynb @@ -0,0 +1,1136 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# spaCy PII Pipeline (Notebook version)\n", + "\n", + "Ноутбук автоматически собран из `spacy_pii_pipeline.py` и разбит на логические секции.\n" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "\nUnified spaCy Pipeline for Russian PII Detection\n=================================================\nCombines a custom regex matcher (EntityRuler-style) with a trained spaCy NER\nin a single pipeline for detecting 30 categories of PII in Russian banking text.\n\nArchitecture:\n text → transformer(BERT) → regex_pii_matcher → ner → entity_merger → output\n\nUsage:\n python spacy_pii_pipeline.py --train # Train the pipeline\n python spacy_pii_pipeline.py --evaluate # Evaluate on dev set\n python spacy_pii_pipeline.py --predict # Predict on test set\n\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "import re\n", + "import csv\n", + "import ast\n", + "import json\n", + "import random\n", + "import argparse\n", + "import logging\n", + "import warnings\n", + "from pathlib import Path\n", + "from copy import deepcopy\n", + "from typing import List, Tuple, Dict, Optional, Set\n", + "from collections import defaultdict\n", + "\n", + "import spacy\n", + "from spacy.language import Language\n", + "from spacy.tokens import Doc, Span\n", + "from spacy.training import Example\n", + "from spacy.util import minibatch, compounding\n", + "\n", + "warnings.filterwarnings(\"ignore\", category=UserWarning)\n", + "logging.basicConfig(level=logging.INFO, format=\"%(asctime)s %(levelname)s %(message)s\")\n", + "logger = logging.getLogger(__name__)\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# LABEL MAPPING (short spaCy label ↔ full competition label)\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "\n", + "LABEL_TO_FULL = {\n", + " \"API_KEY\": \"API ключи\",\n", + " \"CVV\": \"CVV/CVC\",\n", + " \"EMAIL\": \"Email\",\n", + " \"DRIVER_LIC\": \"Водительское удостоверение\",\n", + " \"TEMP_ID\": \"Временное удостоверение личности\",\n", + " \"CITIZENSHIP\": \"Гражданство и названия стран\",\n", + " \"VEHICLE\": \"Данные об автомобиле клиента\",\n", + " \"ORG_DATA\": \"Данные об организации/юридическом лице (ИНН, КПП, ОГРН, БИК, адреса, расчётный счёт)\",\n", + " \"CARD_EXP\": \"Дата окончания срока действия карты\",\n", + " \"REG_DATE\": \"Дата регистрации по месту жительства или пребывания\",\n", + " \"DOB\": \"Дата рождения\",\n", + " \"CARD_HOLDER\": \"Имя держателя карты\",\n", + " \"CODE_WORD\": \"Кодовые слова\",\n", + " \"BIRTHPLACE\": \"Место рождения\",\n", + " \"BANK_NAME\": \"Наименование банка\",\n", + " \"BANK_ACCT\": \"Номер банковского счета\",\n", + " \"CARD_NUM\": \"Номер карты\",\n", + " \"PHONE\": \"Номер телефона\",\n", + " \"OTP\": \"Одноразовые коды\",\n", + " \"PIN\": \"ПИН код\",\n", + " \"PASSWORD\": \"Пароли\",\n", + " \"PASSPORT\": \"Паспортные данные\",\n", + " \"ADDRESS\": \"Полный адрес\",\n", + " \"WORK_PERMIT\": \"Разрешение на работу / визу\",\n", + " \"SNILS\": \"СНИЛС клиента\",\n", + " \"INN\": \"Сведения об ИНН\",\n", + " \"BIRTH_CERT\": \"Свидетельство о рождении\",\n", + " \"RESIDENCE\": \"Серия и номер вида на жительство\",\n", + " \"MAG_STRIPE\": \"Содержимое магнитной полосы\",\n", + " \"FIO\": \"ФИО\",\n", + "}\n", + "\n", + "FULL_TO_LABEL = {v: k for k, v in LABEL_TO_FULL.items()}\n", + "ALL_LABELS = sorted(LABEL_TO_FULL.keys())\n", + "\n", + "# Categories best handled by regex (structured patterns)\n", + "REGEX_LABELS: Set[str] = {\n", + " \"EMAIL\", \"PHONE\", \"CARD_NUM\", \"BANK_ACCT\", \"SNILS\", \"INN\",\n", + " \"CARD_EXP\", \"MAG_STRIPE\", \"API_KEY\", \"CVV\", \"PIN\", \"OTP\",\n", + " \"DRIVER_LIC\", \"TEMP_ID\", \"BIRTH_CERT\", \"RESIDENCE\", \"WORK_PERMIT\",\n", + " \"PASSPORT\", \"PASSWORD\", \"CODE_WORD\", \"ORG_DATA\", \"VEHICLE\",\n", + " \"DOB\", \"REG_DATE\",\n", + "}\n", + "\n", + "# Categories that rely primarily on NER (contextual)\n", + "NER_LABELS: Set[str] = {\n", + " \"FIO\", \"ADDRESS\", \"BIRTHPLACE\", \"CITIZENSHIP\", \"BANK_NAME\",\n", + " \"CARD_HOLDER\",\n", + "}\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# DATA LOADING\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "\n", + "def load_train_data(path: str) -> List[Dict]:\n", + " data = []\n", + " with open(path, \"r\", encoding=\"utf-8\") as f:\n", + " reader = csv.DictReader(f, delimiter=\"\\t\")\n", + " for row in reader:\n", + " text = row[\"text\"]\n", + " target_str = row[\"target\"].strip()\n", + " if target_str == \"[]\":\n", + " entities = []\n", + " else:\n", + " raw = ast.literal_eval(target_str)\n", + " entities = []\n", + " for start, end, full_label in raw:\n", + " label = FULL_TO_LABEL.get(full_label)\n", + " if label:\n", + " entities.append((start, end, label))\n", + " entities = _clean_overlapping(entities)\n", + " data.append({\"text\": text, \"entities\": entities})\n", + " return data\n", + "\n", + "\n", + "def load_test_data(path: str) -> List[Dict]:\n", + " data = []\n", + " with open(path, \"r\", encoding=\"utf-8\") as f:\n", + " reader = csv.DictReader(f)\n", + " for row in reader:\n", + " data.append({\"id\": int(row[\"id\"]), \"text\": row[\"text\"]})\n", + " return data\n", + "\n", + "\n", + "def _clean_overlapping(entities: List[Tuple]) -> List[Tuple]:\n", + " if len(entities) <= 1:\n", + " return entities\n", + " sorted_ents = sorted(entities, key=lambda x: (x[0], -(x[1] - x[0])))\n", + " cleaned = [sorted_ents[0]]\n", + " for ent in sorted_ents[1:]:\n", + " prev = cleaned[-1]\n", + " if ent[0] >= prev[1]:\n", + " cleaned.append(ent)\n", + " elif (ent[1] - ent[0]) > (prev[1] - prev[0]):\n", + " cleaned[-1] = ent\n", + " return cleaned\n", + "\n", + "\n", + "def split_data(data: List[Dict], dev_ratio: float = 0.2, seed: int = 42):\n", + " rng = random.Random(seed)\n", + " shuffled = data.copy()\n", + " rng.shuffle(shuffled)\n", + " n = int(len(shuffled) * (1 - dev_ratio))\n", + " return shuffled[:n], shuffled[n:]\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# REGEX PATTERNS (character-level, context-aware)\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# Each pattern is (compiled_regex, context_regex_or_None, search_radius)\n", + "# context_regex is checked against the full text; if None, pattern applies unconditionally.\n", + "\n", + "_MONTHS_RU = (\n", + " r\"(?:январ[яьие]|феврал[яьие]|март[аеу]?|апрел[яьие]|\"\n", + " r\"ма[яйюе]|июн[яьие]|июл[яьие]|август[аеу]?|\"\n", + " r\"сентябр[яьие]|октябр[яьие]|ноябр[яьие]|декабр[яьие])\"\n", + ")\n", + "\n", + "_DATE_DMY = r\"\\d{2}\\.\\d{2}\\.\\d{4}\"\n", + "_DATE_TEXT = r\"\\d{1,2}\\s\" + _MONTHS_RU + r\"\\s\\d{4}(?:\\s?года?)?\"\n", + "\n", + "\n", + "def _build_regex_rules() -> Dict[str, list]:\n", + " \"\"\"Return {label: [(compiled_pattern, context_pattern|None), ...]}\"\"\"\n", + " rules: Dict[str, list] = {}\n", + "\n", + " # ---- EMAIL ----\n", + " rules[\"EMAIL\"] = [\n", + " (re.compile(r\"[\\w.+-]+@[\\w.-]+\\.\\w{2,}\"), None),\n", + " ]\n", + "\n", + " # ---- PHONE ----\n", + " rules[\"PHONE\"] = [\n", + " (re.compile(r\"(?:\\+7|8)\\s*[\\(\\[]?\\d{3}[\\)\\]]?\\s*\\d{3}[\\-\\s]?\\d{2}[\\-\\s]?\\d{2}\"), None),\n", + " (re.compile(r\"\\b\\d{10,11}\\b\"),\n", + " re.compile(r\"(?i)телефон|звон|привязан|контакт|смс|sms|мобильн\")),\n", + " ]\n", + "\n", + " # ---- CARD_NUM (16 digits, optionally grouped, NOT followed by =) ----\n", + " rules[\"CARD_NUM\"] = [\n", + " (re.compile(r\"\\b\\d{4}[\\s]?\\d{4}[\\s]?\\d{4}[\\s]?\\d{4}\\b(?!=)\"),\n", + " re.compile(r\"(?i)карт|card|оплат|списан|баланс|перевод|блокир\")),\n", + " ]\n", + "\n", + " # ---- BANK_ACCT (20 digits, starts with 4/3) ----\n", + " rules[\"BANK_ACCT\"] = [\n", + " (re.compile(r\"\\b[2-4]\\d{19}\\b\"),\n", + " re.compile(r\"(?i)счёт|счет|account|р/с|расчётн|расчетн\")),\n", + " ]\n", + "\n", + " # ---- SNILS ----\n", + " rules[\"SNILS\"] = [\n", + " (re.compile(r\"\\b\\d{3}-\\d{3}-\\d{3}\\s?\\d{2}\\b\"), None),\n", + " ]\n", + "\n", + " # ---- INN (personal 12 / org 10 digits) ----\n", + " rules[\"INN\"] = [\n", + " (re.compile(r\"\\b\\d{12}\\b\"), re.compile(r\"(?i)инн\")),\n", + " (re.compile(r\"\\b\\d{10}\\b\"), re.compile(r\"(?i)инн\")),\n", + " ]\n", + "\n", + " # ---- CARD_EXP (MM/YY) ----\n", + " rules[\"CARD_EXP\"] = [\n", + " (re.compile(r\"\\b(?:0[1-9]|1[0-2])/\\d{2}\\b\"), None),\n", + " ]\n", + "\n", + " # ---- MAG_STRIPE ----\n", + " rules[\"MAG_STRIPE\"] = [\n", + " (re.compile(r\"%B[\\w\\^/]+\\?\"), None),\n", + " (re.compile(r\"\\b\\d{16}=\\d{15,25}\\b\"), None),\n", + " (re.compile(r\"\\b9F[0-9A-Fa-f]{2,}[0-9A-Fa-f]+\\b\"),\n", + " re.compile(r\"(?i)emv|магнит|track|полос|операци\")),\n", + " ]\n", + "\n", + " # ---- API_KEY ----\n", + " rules[\"API_KEY\"] = [\n", + " (re.compile(r\"\\bAIzaSy[\\w_-]{30,}\"), None),\n", + " (re.compile(r\"\\bGOCSPX-[\\w_-]+\"), None),\n", + " (re.compile(r\"\\bsk_(?:live|test)_[\\w]+\"), None),\n", + " (re.compile(r\"\\bpk_(?:live|test)_[\\w]+\"), None),\n", + " (re.compile(r\"\\bbk_api_key_[\\w]+\"), None),\n", + " (re.compile(r\"\\bdev_key_[\\w]+\"), None),\n", + " (re.compile(r\"\\b\\d{9,10}:[A-Za-z0-9_-]{30,}\"), None),\n", + " (re.compile(r\"\\beyJ[A-Za-z0-9_-]+\\.eyJ[A-Za-z0-9_-]+\\.[A-Za-z0-9_-]+\"), None),\n", + " (re.compile(r\"(?<=['\\\"])[\\w_]+(?=['\\\"])\"),\n", + " re.compile(r\"(?i)ключ|key|api|токен|token\")),\n", + " (re.compile(r\"\\b[A-Z][A-Z0-9_]{3,}(?:_[A-Z0-9]+)+\\b\"),\n", + " re.compile(r\"(?i)ключ|key|api|токен|token|документац|удалени|секрет\")),\n", + " ]\n", + "\n", + " # ---- CVV ----\n", + " rules[\"CVV\"] = [\n", + " (re.compile(r\"\\b\\d{3}\\b\"), re.compile(r\"(?i)cvv|cvc\")),\n", + " ]\n", + "\n", + " # ---- PIN ----\n", + " rules[\"PIN\"] = [\n", + " (re.compile(r\"\\b\\d{4}\\b\"), re.compile(r\"(?i)пин|pin\")),\n", + " ]\n", + "\n", + " # ---- OTP ----\n", + " rules[\"OTP\"] = [\n", + " (re.compile(r\"\\b\\d{4,6}\\b\"),\n", + " re.compile(r\"(?i)(?:одноразов|sms|смс).*код|код.*(?:подтвержд|верифик|вход|оплат|sms|смс)|\"\n", + " r\"пришёл\\s+код|код\\s+(?:\\d|не\\s)\")),\n", + " ]\n", + "\n", + " # ---- DRIVER_LIC ----\n", + " rules[\"DRIVER_LIC\"] = [\n", + " (re.compile(r\"\\b\\d{2}\\s\\d{2}\\s\\d{6}\\b\"),\n", + " re.compile(r\"(?i)водител|удостоверени|ву\\b\")),\n", + " (re.compile(r\"\\b\\d{4}\\s?\\d{6}\\b\"),\n", + " re.compile(r\"(?i)водител|удостоверени|ву\\b\")),\n", + " (re.compile(r\"\\b\\d{10}\\b\"),\n", + " re.compile(r\"(?i)водител|удостоверени\")),\n", + " ]\n", + "\n", + " # ---- TEMP_ID ----\n", + " rules[\"TEMP_ID\"] = [\n", + " (re.compile(r\"\\b[А-ЯA-Z]{2}\\d{8,10}\\b\"),\n", + " re.compile(r\"(?i)временн|удостоверени\")),\n", + " (re.compile(r\"\\b\\d{6,12}\\b\"),\n", + " re.compile(r\"(?i)временн.*(?:удостоверени|документ)|(?:удостоверени|документ).*временн\")),\n", + " ]\n", + "\n", + " # ---- BIRTH_CERT ----\n", + " rules[\"BIRTH_CERT\"] = [\n", + " (re.compile(r\"\\b[IVXLCDM]{1,5}-[А-ЯЁ]{2}\\s\\d{6}\\b\"), None),\n", + " ]\n", + "\n", + " # ---- RESIDENCE (вид на жительство) ----\n", + " rules[\"RESIDENCE\"] = [\n", + " (re.compile(r\"\\b\\d{2}[\\s№-]+\\d{5,7}\\b\"),\n", + " re.compile(r\"(?i)вид на жительств\")),\n", + " ]\n", + "\n", + " # ---- WORK_PERMIT (виза / разрешение на работу) ----\n", + " rules[\"WORK_PERMIT\"] = [\n", + " (re.compile(r\"\\b\\d{2}\\s\\d{7}\\b\"),\n", + " re.compile(r\"(?i)виз[аыуе]|разрешени.*работ\")),\n", + " ]\n", + "\n", + " # ---- PASSPORT ----\n", + " rules[\"PASSPORT\"] = [\n", + " (re.compile(r\"\\b\\d{2}\\s?\\d{2}\\s\\d{6}\\b\"),\n", + " re.compile(r\"(?i)паспорт|серии|серия\")),\n", + " (re.compile(\n", + " r\"(?:[ОоУу](?:правлени|тдел)\\w*\\s+)?(?:ОУФМС|УФМС|ФМС|МВД|ГУВД|ОВД|УВД)\"\n", + " r\"[\\w\\s,.\\-()]*?(?:(?:по|города?|обл(?:асти)?|респ(?:ублик[аие])?|р-на?|района?|край|края|\"\n", + " r\"[А-ЯЁ][а-яё]+(?:ского|ской|ому|ой)?)\\s*)+\"\n", + " ), re.compile(r\"(?i)паспорт|выда\")),\n", + " (re.compile(_DATE_DMY),\n", + " re.compile(r\"(?i)паспорт.*выда|выда.*паспорт|выданн\")),\n", + " (re.compile(_DATE_TEXT),\n", + " re.compile(r\"(?i)паспорт.*выда|выда.*паспорт|выданн\")),\n", + " ]\n", + "\n", + " # ---- PASSWORD ----\n", + " rules[\"PASSWORD\"] = [\n", + " (re.compile(r'(?<=[\"«])[^\\s\"«»]{4,}(?=[\"»])'),\n", + " re.compile(r\"(?i)парол\")),\n", + " (re.compile(r'(?<=пароль\\s)[\"\\s]*\\S{4,}'),\n", + " None),\n", + " (re.compile(r'(?<=пароля\\s)[\"\\s]*\\S{4,}'),\n", + " None),\n", + " (re.compile(r'(?<=например,\\s)\\S{4,}'),\n", + " re.compile(r\"(?i)парол\")),\n", + " ]\n", + "\n", + " # ---- CODE_WORD ----\n", + " rules[\"CODE_WORD\"] = [\n", + " (re.compile(r\"(?<=кодовое слово\\s)[«\\\"]?[А-ЯЁа-яё]{2,}[»\\\"]?\"), None),\n", + " (re.compile(r\"(?<=кодовое слово — )[«\\\"]?[А-ЯЁа-яё]{2,}[»\\\"]?\"), None),\n", + " (re.compile(r\"(?<=кодовое слово — «)[А-ЯЁа-яё]{2,}(?=»)\"), None),\n", + " (re.compile(r\"(?<=слово\\s)«[А-ЯЁа-яё]{2,}»\"),\n", + " re.compile(r\"(?i)кодов\")),\n", + " (re.compile(r\"(?<=«)[А-ЯЁа-яё]{2,}(?=»)\"),\n", + " re.compile(r\"(?i)кодов\\w+\\s+слов\")),\n", + " ]\n", + "\n", + " # ---- DOB ----\n", + " rules[\"DOB\"] = [\n", + " (re.compile(_DATE_DMY + r\"(?:\\s+в\\s+\\d{1,2}:\\d{2}(?::\\d{2})?)?\"),\n", + " re.compile(r\"(?i)рожден|родил\")),\n", + " (re.compile(_DATE_TEXT), re.compile(r\"(?i)рожден|родил\")),\n", + " (re.compile(r\"\\b(?:19|20)\\d{2}\\b\"),\n", + " re.compile(r\"(?i)(?:год|дат)\\w*\\s+рожден|рожден\\w*.*\\b\\d{4}\")),\n", + " (re.compile(_MONTHS_RU, re.I), re.compile(r\"(?i)рожден|родил\")),\n", + " (re.compile(r\"\\b\\d{1,2}\\b\"),\n", + " re.compile(r\"(?i)(?:рожден|родил).*(?:час|минут)|(?:час|минут).*рожден\")),\n", + " ]\n", + "\n", + " # ---- REG_DATE ----\n", + " rules[\"REG_DATE\"] = [\n", + " (re.compile(_DATE_DMY),\n", + " re.compile(r\"(?i)регистрац|пребыван|прописк|местожительств\")),\n", + " (re.compile(_DATE_TEXT),\n", + " re.compile(r\"(?i)регистрац|пребыван|прописк|местожительств\")),\n", + " (re.compile(r\"\\d{1,2}\\s\" + _MONTHS_RU),\n", + " re.compile(r\"(?i)регистрац|пребыван|прописк|местожительств\")),\n", + " ]\n", + "\n", + " # ---- ORG_DATA ----\n", + " rules[\"ORG_DATA\"] = [\n", + " (re.compile(r\"\\b1\\d{12}\\b\"), re.compile(r\"(?i)огрн\")),\n", + " (re.compile(r\"\\b\\d{9}\\b\"), re.compile(r\"(?i)кпп|бик\")),\n", + " (re.compile(r\"\\b[34]\\d{19}\\b\"), re.compile(r\"(?i)расчётн|расчетн|р/с\")),\n", + " (re.compile(r\"\\b\\d{10}\\b\"),\n", + " re.compile(r\"(?i)(?:инн|огрн).*(?:организац|юридич|компан|ооо|зао|пао)|\"\n", + " r\"(?:организац|юридич|компан|ооо|зао|пао).*(?:инн|огрн)\")),\n", + " (re.compile(\n", + " r\"(?:г\\.\\s?[А-ЯЁ][а-яё]+(?:,?\\s*ул\\.\\s*[А-ЯЁа-яё]+(?:,?\\s*д\\.\\s*\\d+)?)?)\"\n", + " ), re.compile(r\"(?i)юридич.*адрес|адрес.*юридич|организац\")),\n", + " ]\n", + "\n", + " # ---- VEHICLE ----\n", + " rules[\"VEHICLE\"] = [\n", + " (re.compile(r\"\\b[A-Z0-9]{17}\\b\"),\n", + " re.compile(r\"(?i)vin|автомобил|машин|авто(?:кредит|страхов)\")),\n", + " (re.compile(r\"\\b(?:19|20)\\d{2}\\b\"),\n", + " re.compile(r\"(?i)(?:авто|машин|vin).*(?:год|г\\.)|\\bгод[а]?\\s+(?:выпуск|(?:19|20)\\d{2})\")),\n", + " ]\n", + "\n", + " return rules\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# REGEX DETECTION ENGINE\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "\n", + "class RegexPIIDetector:\n", + " def __init__(self):\n", + " self.rules = _build_regex_rules()\n", + "\n", + " def detect(self, text: str) -> List[Tuple[int, int, str]]:\n", + " all_matches = []\n", + " for label, pattern_list in self.rules.items():\n", + " for pat, ctx in pattern_list:\n", + " if ctx is not None and not ctx.search(text):\n", + " continue\n", + " for m in pat.finditer(text):\n", + " s, e = m.start(), m.end()\n", + " if e - s < 1:\n", + " continue\n", + " all_matches.append((s, e, label))\n", + " return _resolve_overlaps(all_matches)\n", + "\n", + "\n", + "def _resolve_overlaps(matches: List[Tuple[int, int, str]]) -> List[Tuple[int, int, str]]:\n", + " if not matches:\n", + " return []\n", + " sorted_m = sorted(matches, key=lambda x: (x[0], -(x[1] - x[0])))\n", + " result = [sorted_m[0]]\n", + " for m in sorted_m[1:]:\n", + " prev = result[-1]\n", + " if m[0] >= prev[1]:\n", + " result.append(m)\n", + " elif (m[1] - m[0]) > (prev[1] - prev[0]):\n", + " result[-1] = m\n", + " return result\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# CUSTOM SPACY COMPONENTS\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "\n", + "@Language.factory(\"regex_pii_matcher\")\n", + "class RegexPIIMatcher:\n", + " \"\"\"spaCy component: finds PII via regex and stores in doc.spans['regex'].\"\"\"\n", + "\n", + " def __init__(self, nlp, name):\n", + " self.detector = RegexPIIDetector()\n", + "\n", + " def __call__(self, doc: Doc) -> Doc:\n", + " text = doc.text\n", + " matches = self.detector.detect(text)\n", + " spans = []\n", + " for start, end, label in matches:\n", + " span = doc.char_span(start, end, label=label, alignment_mode=\"expand\")\n", + " if span is not None:\n", + " spans.append(span)\n", + " doc.spans[\"regex\"] = spans\n", + " return doc\n", + "\n", + "\n", + "_QUOTE_CHARS = set('\"«»\\'\"\\'')\n", + "_REG_DATE_PREFIX = re.compile(r\"^(?:с|до|от|по|на|)\\s+\", re.I)\n", + "_ORG_CONTEXT = re.compile(r\"(?i)расчётн|расчетн|р/с|организац|юридич|компан|ооо|зао|пао|ао\\b\")\n", + "\n", + "\n", + "@Language.factory(\"entity_merger\")\n", + "class EntityMerger:\n", + " \"\"\"Merges doc.spans['regex'] with doc.ents, prioritizing regex matches.\n", + " Also applies post-processing fixes for boundary issues.\n", + " Stores exact char-level results in doc.user_data['pii_entities'].\"\"\"\n", + "\n", + " def __init__(self, nlp, name):\n", + " pass\n", + "\n", + " def __call__(self, doc: Doc) -> Doc:\n", + " text = doc.text\n", + " regex_spans = list(doc.spans.get(\"regex\", []))\n", + " ner_ents = list(doc.ents)\n", + "\n", + " all_ents = []\n", + " for span in regex_spans:\n", + " all_ents.append((span.start_char, span.end_char, span.label_, \"regex\"))\n", + " for ent in ner_ents:\n", + " all_ents.append((ent.start_char, ent.end_char, ent.label_, \"ner\"))\n", + "\n", + " merged = _merge_entities(all_ents)\n", + " merged = _postprocess_entities(text, merged)\n", + "\n", + " doc.user_data[\"pii_entities\"] = [(s, e, lbl) for s, e, lbl, _ in merged]\n", + "\n", + " final_spans = []\n", + " for start, end, label, _ in merged:\n", + " span = doc.char_span(start, end, label=label, alignment_mode=\"expand\")\n", + " if span is not None and len(span) > 0:\n", + " final_spans.append(span)\n", + "\n", + " try:\n", + " doc.ents = final_spans\n", + " except ValueError:\n", + " doc.ents = _deduplicate_spans(final_spans)\n", + " return doc\n", + "\n", + "\n", + "_ORG_KEYWORDS = re.compile(\n", + " r\"(?i)организаци|юридич|компани|ооо\\b|зао\\b|пао\\b|оао\\b|ао\\s*[\\\"«]|\"\n", + " r\"контрагент|реквизит|фирм|предприяти\"\n", + ")\n", + "\n", + "_ORG_DATA_MARKERS = re.compile(\n", + " r\"(?i)кпп|огрн|бик\\b|сменил\\w*\\s+адрес|юридическ\\w*\\s+адрес|\"\n", + " r\"почтов\\w*\\s+адрес|данные?\\s+по\\s+инн|данные?\\s+для\\b|\"\n", + " r\"актуальн\\w*\\s+инн|инн\\s+актуальн|инн.*указан.*неверн|\"\n", + " r\"для\\s+организаци\\w*\\s+с\\s+инн|контрагент\\w*\\s+по\\s+инн\"\n", + ")\n", + "\n", + "\n", + "def _postprocess_entities(text: str, ents: List[Tuple]) -> List[Tuple]:\n", + " \"\"\"Fix common boundary issues after merge.\"\"\"\n", + " result = []\n", + " text_lower = text.lower()\n", + " has_org_context = bool(_ORG_KEYWORDS.search(text))\n", + " is_org_data_text = bool(_ORG_DATA_MARKERS.search(text))\n", + "\n", + " for start, end, label, source in ents:\n", + " s, e = start, end\n", + "\n", + " # Strip quotes from PASSWORD and CODE_WORD\n", + " if label in (\"PASSWORD\", \"CODE_WORD\"):\n", + " while s < e and text[s] in _QUOTE_CHARS:\n", + " s += 1\n", + " while e > s and text[e - 1] in _QUOTE_CHARS:\n", + " e -= 1\n", + " while e > s and text[e - 1] in (',', '.', ' '):\n", + " e -= 1\n", + "\n", + " # BANK_ACCT → ORG_DATA\n", + " if label == \"BANK_ACCT\":\n", + " before = text[max(0, s - 40):s].lower()\n", + " if re.search(r\"расчётн|расчетн|р/с\", before):\n", + " label = \"ORG_DATA\"\n", + " elif re.search(r\"реквизит\", before) and has_org_context:\n", + " label = \"ORG_DATA\"\n", + " elif is_org_data_text:\n", + " label = \"ORG_DATA\"\n", + "\n", + " # INN → ORG_DATA\n", + " if label == \"INN\":\n", + " if is_org_data_text:\n", + " label = \"ORG_DATA\"\n", + "\n", + " # CARD_EXP: strip \"года\" suffix if present (inconsistent in gold)\n", + " if label == \"CARD_EXP\":\n", + " chunk = text[s:e]\n", + " m_year_suffix = re.search(r\"\\s+года?$\", chunk)\n", + " if m_year_suffix:\n", + " e = s + m_year_suffix.start()\n", + "\n", + " # DRIVER_LIC ↔ PASSPORT disambiguation based on context\n", + " if label in (\"DRIVER_LIC\", \"PASSPORT\"):\n", + " full_lower = text_lower\n", + " has_passport_ctx = bool(re.search(\n", + " r\"паспорт|серии\\b|серия\\b|выдан|паспортн\", full_lower\n", + " ))\n", + " has_driver_ctx = bool(re.search(\n", + " r\"водител|удостоверени|\\bву\\b|страхов|каско|осаго|\"\n", + " r\"автомобил|авто\\b|транспорт|дтп|штраф.*гибдд|гибдд\",\n", + " full_lower\n", + " ))\n", + " if label == \"DRIVER_LIC\" and has_passport_ctx and not has_driver_ctx:\n", + " label = \"PASSPORT\"\n", + " elif label == \"PASSPORT\" and has_driver_ctx and not has_passport_ctx:\n", + " label = \"DRIVER_LIC\"\n", + "\n", + " # Strip trailing punctuation from API_KEY\n", + " if label == \"API_KEY\":\n", + " while e > s and text[e - 1] in ('.', ',', ';', ':', '!', '?', ' '):\n", + " e -= 1\n", + "\n", + " # DOB: split full text dates into parts (day, month, year)\n", + " # because gold annotation is split ~50% of the time\n", + " if label == \"DOB\":\n", + " chunk = text[s:e]\n", + " m_text_date = re.match(\n", + " r'^(\\d{1,2})\\s+(' + _MONTHS_RU + r')\\s+(\\d{4})(?:\\s+года?)?$',\n", + " chunk, re.I\n", + " )\n", + " if m_text_date:\n", + " day_s = s + m_text_date.start(1)\n", + " day_e = s + m_text_date.end(1)\n", + " mon_s = s + m_text_date.start(2)\n", + " mon_e = s + m_text_date.end(2)\n", + " yr_s = s + m_text_date.start(3)\n", + " yr_e = s + m_text_date.end(3)\n", + " result.append((day_s, day_e, \"DOB\", source))\n", + " result.append((mon_s, mon_e, \"DOB\", source))\n", + " result.append((yr_s, yr_e, \"DOB\", source))\n", + " continue\n", + "\n", + " if s < e:\n", + " result.append((s, e, label, source))\n", + " return result\n", + "\n", + "\n", + "def _merge_entities(ents: List[Tuple]) -> List[Tuple]:\n", + " \"\"\"Merge entities from regex and NER.\n", + " When overlapping: prefer longer span; if same length prefer regex.\"\"\"\n", + " if not ents:\n", + " return []\n", + " sorted_e = sorted(ents, key=lambda x: (x[0], -(x[1] - x[0])))\n", + " result = [sorted_e[0]]\n", + " for e in sorted_e[1:]:\n", + " prev = result[-1]\n", + " if e[0] >= prev[1]:\n", + " result.append(e)\n", + " continue\n", + " prev_len = prev[1] - prev[0]\n", + " cur_len = e[1] - e[0]\n", + " if cur_len > prev_len:\n", + " result[-1] = e\n", + " elif cur_len == prev_len and e[3] == \"regex\" and prev[3] == \"ner\":\n", + " result[-1] = e\n", + " elif e[3] == \"regex\" and prev[3] == \"ner\" and cur_len >= prev_len * 0.8:\n", + " result[-1] = e\n", + " return result\n", + "\n", + "\n", + "def _deduplicate_spans(spans: List[Span]) -> List[Span]:\n", + " \"\"\"Remove overlapping spans keeping the first one.\"\"\"\n", + " if not spans:\n", + " return []\n", + " sorted_s = sorted(spans, key=lambda s: (s.start, -len(s)))\n", + " result = [sorted_s[0]]\n", + " for s in sorted_s[1:]:\n", + " if s.start >= result[-1].end:\n", + " result.append(s)\n", + " return result\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# TRAINING DATA PREPARATION\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "\n", + "def prepare_examples(nlp, data: List[Dict]) -> List[Example]:\n", + " \"\"\"Convert dataset to spaCy Example objects for NER training.\"\"\"\n", + " examples = []\n", + " skipped = 0\n", + " for item in data:\n", + " text = item[\"text\"]\n", + " entities = item[\"entities\"]\n", + " doc = nlp.make_doc(text)\n", + " ents_dict = {\"entities\": [(s, e, lbl) for s, e, lbl in entities]}\n", + " try:\n", + " example = Example.from_dict(doc, ents_dict)\n", + " examples.append(example)\n", + " except Exception:\n", + " skipped += 1\n", + " if skipped:\n", + " logger.warning(\"Skipped %d examples with alignment issues\", skipped)\n", + " return examples\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# PIPELINE BUILD & TRAINING\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "\n", + "def build_pipeline() -> Language:\n", + " \"\"\"Build the unified spaCy pipeline with BERT-backed transformer.\"\"\"\n", + " logger.info(\"Loading spaCy ru tokenizer and adding BERT transformer...\")\n", + " nlp = spacy.blank(\"ru\")\n", + "\n", + " nlp.add_pipe(\n", + " \"transformer\",\n", + " config={\n", + " \"model\": {\n", + " \"@architectures\": \"spacy-transformers.TransformerModel.v3\",\n", + " \"name\": \"DeepPavlov/rubert-base-cased\",\n", + " \"tokenizer_config\": {\"use_fast\": True},\n", + " }\n", + " },\n", + " )\n", + " nlp.add_pipe(\"regex_pii_matcher\", after=\"transformer\")\n", + " ner = nlp.add_pipe(\"ner\", config=_build_transformer_ner_config(), last=True)\n", + " nlp.add_pipe(\"entity_merger\", last=True)\n", + "\n", + " for label in ALL_LABELS:\n", + " ner.add_label(label)\n", + "\n", + " return nlp\n", + "\n", + "\n", + "def _build_transformer_ner_config(hidden_width: int = 64) -> Dict:\n", + " \"\"\"NER head config that listens to BERT transformer embeddings.\"\"\"\n", + " return {\n", + " \"model\": {\n", + " \"@architectures\": \"spacy.TransitionBasedParser.v2\",\n", + " \"state_type\": \"ner\",\n", + " \"extra_state_tokens\": False,\n", + " \"hidden_width\": hidden_width,\n", + " \"maxout_pieces\": 2,\n", + " \"use_upper\": True,\n", + " \"nO\": None,\n", + " \"tok2vec\": {\n", + " \"@architectures\": \"spacy-transformers.TransformerListener.v1\",\n", + " \"upstream\": \"transformer\",\n", + " \"grad_factor\": 1.0,\n", + " \"pooling\": {\"@layers\": \"reduce_mean.v1\"},\n", + " },\n", + " }\n", + " }\n", + "\n", + "\n", + "def train_pipeline(\n", + " nlp: Language,\n", + " train_data: List[Dict],\n", + " dev_data: List[Dict],\n", + " n_epochs: int = 30,\n", + " batch_size_start: float = 4.0,\n", + " batch_size_end: float = 32.0,\n", + " drop: float = 0.35,\n", + " patience: int = 5,\n", + " output_dir: str = \"pii_spacy_model\",\n", + "):\n", + " \"\"\"Train the NER component of the pipeline.\"\"\"\n", + " logger.info(\"Preparing training examples... (%d train, %d dev)\", len(train_data), len(dev_data))\n", + "\n", + " train_examples = prepare_examples(nlp, train_data)\n", + "\n", + " logger.info(\"Training examples: %d, Dev items: %d\", len(train_examples), len(dev_data))\n", + "\n", + " get_examples = lambda: train_examples[:200]\n", + "\n", + " with nlp.select_pipes(enable=[\"transformer\", \"ner\"]):\n", + " nlp.initialize(get_examples)\n", + "\n", + " optimizer = nlp.create_optimizer()\n", + " best_f1 = 0.0\n", + " no_improve = 0\n", + "\n", + " for epoch in range(n_epochs):\n", + " random.shuffle(train_examples)\n", + " losses = {}\n", + " batches = minibatch(train_examples, size=compounding(batch_size_start, batch_size_end, 1.001))\n", + "\n", + " with nlp.select_pipes(enable=[\"transformer\", \"ner\"]):\n", + " for batch in batches:\n", + " nlp.update(batch, sgd=optimizer, losses=losses, drop=drop)\n", + "\n", + " dev_scores = evaluate_on_data(nlp, dev_data)\n", + " f1 = dev_scores[\"f1\"]\n", + " logger.info(\n", + " \"Epoch %02d | loss=%.2f | P=%.3f R=%.3f F1=%.3f\",\n", + " epoch, losses.get(\"ner\", 0), dev_scores[\"precision\"], dev_scores[\"recall\"], f1,\n", + " )\n", + "\n", + " if f1 > best_f1:\n", + " best_f1 = f1\n", + " no_improve = 0\n", + " Path(output_dir).mkdir(exist_ok=True)\n", + " nlp.to_disk(output_dir)\n", + " logger.info(\" → New best F1=%.3f, model saved to %s\", f1, output_dir)\n", + " else:\n", + " no_improve += 1\n", + " if no_improve >= patience:\n", + " logger.info(\"Early stopping at epoch %d (no improvement for %d epochs)\", epoch, patience)\n", + " break\n", + "\n", + " logger.info(\"Training complete. Best F1=%.3f\", best_f1)\n", + " return nlp\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# EVALUATION (strict micro-averaged F1)\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "\n", + "def extract_entities_from_doc(doc: Doc) -> List[Tuple[int, int, str]]:\n", + " \"\"\"Extract (start_char, end_char, label) from a processed doc.\n", + " Uses exact char-level data stored by EntityMerger when available.\"\"\"\n", + " if \"pii_entities\" in doc.user_data:\n", + " return doc.user_data[\"pii_entities\"]\n", + " return [(ent.start_char, ent.end_char, ent.label_) for ent in doc.ents]\n", + "\n", + "\n", + "def evaluate_examples(nlp, examples: List[Example]) -> Dict[str, float]:\n", + " \"\"\"Evaluate on spaCy Example objects.\"\"\"\n", + " tp = fp = fn = 0\n", + " for example in examples:\n", + " gold_ents = set()\n", + " for ent in example.reference.ents:\n", + " gold_ents.add((ent.start_char, ent.end_char, ent.label_))\n", + " if not gold_ents:\n", + " for s, e, lbl in example.y.user_data.get(\"entities\", []):\n", + " gold_ents.add((s, e, lbl))\n", + "\n", + " gold_from_data = set()\n", + " try:\n", + " text = example.reference.text\n", + " ents_in_ref = [(ent.start_char, ent.end_char, ent.label_) for ent in example.reference.ents]\n", + " gold_from_data = set(ents_in_ref) if ents_in_ref else gold_ents\n", + " except Exception:\n", + " gold_from_data = gold_ents\n", + "\n", + " pred_doc = nlp(example.reference.text)\n", + " pred_ents = set(extract_entities_from_doc(pred_doc))\n", + "\n", + " tp += len(pred_ents & gold_from_data)\n", + " fp += len(pred_ents - gold_from_data)\n", + " fn += len(gold_from_data - pred_ents)\n", + "\n", + " precision = tp / (tp + fp) if (tp + fp) > 0 else 0.0\n", + " recall = tp / (tp + fn) if (tp + fn) > 0 else 0.0\n", + " f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0.0\n", + " return {\"precision\": precision, \"recall\": recall, \"f1\": f1, \"tp\": tp, \"fp\": fp, \"fn\": fn}\n", + "\n", + "\n", + "def evaluate_on_data(nlp, data: List[Dict]) -> Dict[str, float]:\n", + " \"\"\"Evaluate on raw data dicts (with char offsets).\"\"\"\n", + " tp = fp = fn = 0\n", + " per_label_tp = defaultdict(int)\n", + " per_label_fp = defaultdict(int)\n", + " per_label_fn = defaultdict(int)\n", + "\n", + " for item in data:\n", + " text = item[\"text\"]\n", + " gold = set((s, e, lbl) for s, e, lbl in item[\"entities\"])\n", + " doc = nlp(text)\n", + " pred = set(extract_entities_from_doc(doc))\n", + "\n", + " matched = pred & gold\n", + " tp += len(matched)\n", + " fp += len(pred - gold)\n", + " fn += len(gold - pred)\n", + "\n", + " for s, e, lbl in matched:\n", + " per_label_tp[lbl] += 1\n", + " for s, e, lbl in (pred - gold):\n", + " per_label_fp[lbl] += 1\n", + " for s, e, lbl in (gold - pred):\n", + " per_label_fn[lbl] += 1\n", + "\n", + " precision = tp / (tp + fp) if (tp + fp) > 0 else 0.0\n", + " recall = tp / (tp + fn) if (tp + fn) > 0 else 0.0\n", + " f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0.0\n", + "\n", + " return {\n", + " \"precision\": precision,\n", + " \"recall\": recall,\n", + " \"f1\": f1,\n", + " \"tp\": tp, \"fp\": fp, \"fn\": fn,\n", + " \"per_label_tp\": dict(per_label_tp),\n", + " \"per_label_fp\": dict(per_label_fp),\n", + " \"per_label_fn\": dict(per_label_fn),\n", + " }\n", + "\n", + "\n", + "def print_per_label_metrics(scores: Dict):\n", + " \"\"\"Pretty-print per-label P/R/F1.\"\"\"\n", + " all_labels = set()\n", + " all_labels.update(scores.get(\"per_label_tp\", {}).keys())\n", + " all_labels.update(scores.get(\"per_label_fp\", {}).keys())\n", + " all_labels.update(scores.get(\"per_label_fn\", {}).keys())\n", + "\n", + " print(f\"\\n{'Label':<16} {'P':>6} {'R':>6} {'F1':>6} {'TP':>5} {'FP':>5} {'FN':>5} Full Name\")\n", + " print(\"-\" * 90)\n", + " for lbl in sorted(all_labels):\n", + " t = scores[\"per_label_tp\"].get(lbl, 0)\n", + " f = scores[\"per_label_fp\"].get(lbl, 0)\n", + " n = scores[\"per_label_fn\"].get(lbl, 0)\n", + " p = t / (t + f) if (t + f) > 0 else 0\n", + " r = t / (t + n) if (t + n) > 0 else 0\n", + " f1 = 2 * p * r / (p + r) if (p + r) > 0 else 0\n", + " full = LABEL_TO_FULL.get(lbl, lbl)\n", + " print(f\"{lbl:<16} {p:>6.3f} {r:>6.3f} {f1:>6.3f} {t:>5d} {f:>5d} {n:>5d} {full}\")\n", + "\n", + " print(\"-\" * 90)\n", + " print(f\"{'MICRO':.<16} {scores['precision']:>6.3f} {scores['recall']:>6.3f} {scores['f1']:>6.3f} \"\n", + " f\"{scores['tp']:>5d} {scores['fp']:>5d} {scores['fn']:>5d}\")\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# INFERENCE\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "\n", + "def predict_text(nlp, text: str) -> List[Tuple[int, int, str]]:\n", + " \"\"\"Predict PII entities and return in competition format (full labels).\n", + " Uses exact char-level positions from EntityMerger.\"\"\"\n", + " doc = nlp(text)\n", + " entities = extract_entities_from_doc(doc)\n", + " result = []\n", + " for s, e, lbl in entities:\n", + " full_label = LABEL_TO_FULL.get(lbl, lbl)\n", + " result.append((s, e, full_label))\n", + " return result\n", + "\n", + "\n", + "def predict_test_set(nlp, test_path: str, output_path: str):\n", + " \"\"\"Run inference on the private test set and save results.\"\"\"\n", + " test_data = load_test_data(test_path)\n", + " results = []\n", + " for item in test_data:\n", + " preds = predict_text(nlp, item[\"text\"])\n", + " results.append({\n", + " \"id\": item[\"id\"],\n", + " \"prediction\": preds,\n", + " })\n", + "\n", + " with open(output_path, \"w\", encoding=\"utf-8\") as f:\n", + " json.dump(results, f, ensure_ascii=False, indent=2)\n", + " logger.info(\"Predictions saved to %s (%d rows)\", output_path, len(results))\n", + " return results\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "# MAIN\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# =====================================================================\n", + "\n", + "def main():\n", + " parser = argparse.ArgumentParser(description=\"spaCy PII Pipeline\")\n", + " parser.add_argument(\"--train\", action=\"store_true\", help=\"Train the pipeline\")\n", + " parser.add_argument(\"--evaluate\", action=\"store_true\", help=\"Evaluate on dev set\")\n", + " parser.add_argument(\"--predict\", action=\"store_true\", help=\"Predict on test set\")\n", + " parser.add_argument(\"--regex-only\", action=\"store_true\", help=\"Evaluate regex-only (no NER)\")\n", + " parser.add_argument(\"--data\", default=\"train_dataset.tsv\", help=\"Training data path\")\n", + " parser.add_argument(\"--test\", default=\"private_test_dataset.csv\", help=\"Test data path\")\n", + " parser.add_argument(\"--model-dir\", default=\"pii_spacy_model\", help=\"Model directory\")\n", + " parser.add_argument(\"--output\", default=\"predictions.json\", help=\"Predictions output path\")\n", + " parser.add_argument(\"--epochs\", type=int, default=30, help=\"Training epochs\")\n", + " parser.add_argument(\"--dev-ratio\", type=float, default=0.2, help=\"Dev split ratio\")\n", + " parser.add_argument(\"--full-train\", action=\"store_true\",\n", + " help=\"Train on full dataset (no dev split) for final submission\")\n", + " args = parser.parse_args()\n", + "\n", + " if args.regex_only:\n", + " logger.info(\"=== Regex-only evaluation ===\")\n", + " data = load_train_data(args.data)\n", + " _, dev_data = split_data(data, args.dev_ratio)\n", + " detector = RegexPIIDetector()\n", + " tp = fp = fn = 0\n", + " per_label_tp = defaultdict(int)\n", + " per_label_fp = defaultdict(int)\n", + " per_label_fn = defaultdict(int)\n", + " for item in dev_data:\n", + " gold = set((s, e, lbl) for s, e, lbl in item[\"entities\"])\n", + " pred = set(tuple(x) for x in detector.detect(item[\"text\"]))\n", + " matched = pred & gold\n", + " tp += len(matched)\n", + " fp += len(pred - gold)\n", + " fn += len(gold - pred)\n", + " for s, e, lbl in matched:\n", + " per_label_tp[lbl] += 1\n", + " for s, e, lbl in (pred - gold):\n", + " per_label_fp[lbl] += 1\n", + " for s, e, lbl in (gold - pred):\n", + " per_label_fn[lbl] += 1\n", + " p = tp / (tp + fp) if (tp + fp) > 0 else 0\n", + " r = tp / (tp + fn) if (tp + fn) > 0 else 0\n", + " f1 = 2 * p * r / (p + r) if (p + r) > 0 else 0\n", + " scores = {\n", + " \"precision\": p, \"recall\": r, \"f1\": f1,\n", + " \"tp\": tp, \"fp\": fp, \"fn\": fn,\n", + " \"per_label_tp\": dict(per_label_tp),\n", + " \"per_label_fp\": dict(per_label_fp),\n", + " \"per_label_fn\": dict(per_label_fn),\n", + " }\n", + " print_per_label_metrics(scores)\n", + " return\n", + "\n", + " if args.full_train:\n", + " data = load_train_data(args.data)\n", + " nlp = build_pipeline()\n", + " small_dev = data[:200]\n", + " nlp = train_pipeline(\n", + " nlp, data, small_dev, n_epochs=args.epochs,\n", + " output_dir=args.model_dir, patience=8,\n", + " )\n", + " logger.info(\"Full-train complete. Model saved to %s\", args.model_dir)\n", + "\n", + " elif args.train:\n", + " data = load_train_data(args.data)\n", + " train_data, dev_data = split_data(data, args.dev_ratio)\n", + " nlp = build_pipeline()\n", + " nlp = train_pipeline(nlp, train_data, dev_data, n_epochs=args.epochs, output_dir=args.model_dir)\n", + " logger.info(\"=== Final evaluation on dev set ===\")\n", + " nlp_best = spacy.load(args.model_dir)\n", + " scores = evaluate_on_data(nlp_best, dev_data)\n", + " print_per_label_metrics(scores)\n", + "\n", + " if args.evaluate:\n", + " data = load_train_data(args.data)\n", + " _, dev_data = split_data(data, args.dev_ratio)\n", + " nlp = spacy.load(args.model_dir)\n", + " scores = evaluate_on_data(nlp, dev_data)\n", + " print_per_label_metrics(scores)\n", + "\n", + " if args.predict:\n", + " nlp = spacy.load(args.model_dir)\n", + " predict_test_set(nlp, args.test, args.output)\n", + "\n", + "\n", + "if __name__ == \"__main__\":\n", + " main()\n" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "name": "python", + "version": "3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} \ No newline at end of file diff --git a/spacy_pii_pipeline.py b/spacy_pii_pipeline.py index 9f4720f..8b7cbe4 100644 --- a/spacy_pii_pipeline.py +++ b/spacy_pii_pipeline.py @@ -5,7 +5,7 @@ in a single pipeline for detecting 30 categories of PII in Russian banking text. Architecture: - text → tok2vec → regex_pii_matcher → ner → entity_merger → output + text → transformer(BERT) → regex_pii_matcher → ner → entity_merger → output Usage: python spacy_pii_pipeline.py --train # Train the pipeline @@ -32,12 +32,6 @@ from spacy.training import Example from spacy.util import minibatch, compounding -try: - import optuna - OPTUNA_AVAILABLE = True -except ImportError: - OPTUNA_AVAILABLE = False - warnings.filterwarnings("ignore", category=UserWarning) logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") logger = logging.getLogger(__name__) @@ -639,12 +633,22 @@ def prepare_examples(nlp, data: List[Dict]) -> List[Example]: # ===================================================================== def build_pipeline() -> Language: - """Build the unified spaCy pipeline.""" - logger.info("Loading ru_core_news_lg base model...") - nlp = spacy.load("ru_core_news_lg", exclude=["ner"]) - - nlp.add_pipe("regex_pii_matcher", before="tok2vec") - ner = nlp.add_pipe("ner", last=True) + """Build the unified spaCy pipeline with BERT-backed transformer.""" + logger.info("Loading spaCy ru tokenizer and adding BERT transformer...") + nlp = spacy.blank("ru") + + nlp.add_pipe( + "transformer", + config={ + "model": { + "@architectures": "spacy-transformers.TransformerModel.v3", + "name": "DeepPavlov/rubert-base-cased", + "tokenizer_config": {"use_fast": True}, + } + }, + ) + nlp.add_pipe("regex_pii_matcher", after="transformer") + ner = nlp.add_pipe("ner", config=_build_transformer_ner_config(), last=True) nlp.add_pipe("entity_merger", last=True) for label in ALL_LABELS: @@ -653,6 +657,27 @@ def build_pipeline() -> Language: return nlp +def _build_transformer_ner_config(hidden_width: int = 64) -> Dict: + """NER head config that listens to BERT transformer embeddings.""" + return { + "model": { + "@architectures": "spacy.TransitionBasedParser.v2", + "state_type": "ner", + "extra_state_tokens": False, + "hidden_width": hidden_width, + "maxout_pieces": 2, + "use_upper": True, + "nO": None, + "tok2vec": { + "@architectures": "spacy-transformers.TransformerListener.v1", + "upstream": "transformer", + "grad_factor": 1.0, + "pooling": {"@layers": "reduce_mean.v1"}, + }, + } + } + + def train_pipeline( nlp: Language, train_data: List[Dict], @@ -667,19 +692,15 @@ def train_pipeline( """Train the NER component of the pipeline.""" logger.info("Preparing training examples... (%d train, %d dev)", len(train_data), len(dev_data)) - tok2vec_bytes = nlp.get_pipe("tok2vec").to_bytes() - train_examples = prepare_examples(nlp, train_data) logger.info("Training examples: %d, Dev items: %d", len(train_examples), len(dev_data)) get_examples = lambda: train_examples[:200] - with nlp.select_pipes(enable=["tok2vec", "ner"]): + with nlp.select_pipes(enable=["transformer", "ner"]): nlp.initialize(get_examples) - nlp.get_pipe("tok2vec").from_bytes(tok2vec_bytes) - optimizer = nlp.create_optimizer() best_f1 = 0.0 no_improve = 0 @@ -689,7 +710,7 @@ def train_pipeline( losses = {} batches = minibatch(train_examples, size=compounding(batch_size_start, batch_size_end, 1.001)) - with nlp.select_pipes(enable=["tok2vec", "ner"]): + with nlp.select_pipes(enable=["transformer", "ner"]): for batch in batches: nlp.update(batch, sgd=optimizer, losses=losses, drop=drop) @@ -716,245 +737,6 @@ def train_pipeline( return nlp -# ===================================================================== -# HYPERPARAMETER TUNING (Optuna) -# ===================================================================== - -def _build_pipeline_with_params( - hidden_width: int = 64, - ner_depth: int = 4, - ner_width: int = 96, - ner_maxout: int = 3, - ner_embed_size: int = 2000, - ner_window_size: int = 1, -) -> Language: - """Build pipeline with custom NER architecture hyperparameters.""" - logger.info("Loading ru_core_news_lg base model...") - nlp = spacy.load("ru_core_news_lg", exclude=["ner"]) - - nlp.add_pipe("regex_pii_matcher", before="tok2vec") - - ner_config = { - "model": { - "@architectures": "spacy.TransitionBasedParser.v2", - "state_type": "ner", - "extra_state_tokens": False, - "hidden_width": hidden_width, - "maxout_pieces": 2, - "use_upper": True, - "nO": None, - "tok2vec": { - "@architectures": "spacy.HashEmbedCNN.v2", - "pretrained_vectors": None, - "width": ner_width, - "depth": ner_depth, - "embed_size": ner_embed_size, - "window_size": ner_window_size, - "maxout_pieces": ner_maxout, - "subword_features": True, - }, - } - } - ner = nlp.add_pipe("ner", config=ner_config, last=True) - nlp.add_pipe("entity_merger", last=True) - - for label in ALL_LABELS: - ner.add_label(label) - - return nlp - - -def _train_for_trial( - nlp: Language, - train_examples: List[Example], - dev_data: List[Dict], - n_epochs: int, - drop: float, - batch_start: float, - batch_end: float, - batch_compound: float, - learn_rate: float, - patience: int = 5, - trial=None, -) -> float: - """Train NER and return best dev F1. Supports Optuna pruning.""" - tok2vec_bytes = nlp.get_pipe("tok2vec").to_bytes() - - get_examples = lambda: train_examples[:200] - with nlp.select_pipes(enable=["tok2vec", "ner"]): - nlp.initialize(get_examples) - - nlp.get_pipe("tok2vec").from_bytes(tok2vec_bytes) - - optimizer = nlp.create_optimizer() - optimizer.learn_rate = learn_rate - - best_f1 = 0.0 - no_improve = 0 - - for epoch in range(n_epochs): - random.shuffle(train_examples) - losses = {} - batches = minibatch( - train_examples, - size=compounding(batch_start, batch_end, batch_compound), - ) - - with nlp.select_pipes(enable=["tok2vec", "ner"]): - for batch in batches: - nlp.update(batch, sgd=optimizer, losses=losses, drop=drop) - - dev_scores = evaluate_on_data(nlp, dev_data) - f1 = dev_scores["f1"] - logger.info( - " [trial] Epoch %02d | loss=%.2f | P=%.3f R=%.3f F1=%.3f", - epoch, losses.get("ner", 0), - dev_scores["precision"], dev_scores["recall"], f1, - ) - - if trial is not None: - trial.report(f1, epoch) - if trial.should_prune(): - raise optuna.TrialPruned() - - if f1 > best_f1: - best_f1 = f1 - no_improve = 0 - else: - no_improve += 1 - if no_improve >= patience: - logger.info(" [trial] Early stopping at epoch %d", epoch) - break - - return best_f1 - - -def run_hyperparameter_tuning( - data_path: str, - dev_ratio: float = 0.2, - n_trials: int = 20, - tuning_epochs: int = 15, - study_name: str = "pii_ner_v5", - storage: Optional[str] = None, -) -> Dict: - """Run Optuna hyperparameter search and return the best params.""" - if not OPTUNA_AVAILABLE: - raise RuntimeError("optuna is not installed. Run: pip install optuna") - - data = load_train_data(data_path) - train_data, dev_data = split_data(data, dev_ratio) - - base_nlp = spacy.load("ru_core_news_lg", exclude=["ner"]) - base_examples = prepare_examples(base_nlp, train_data) - logger.info("Prepared %d training examples for tuning", len(base_examples)) - - def objective(trial: optuna.Trial) -> float: - drop = trial.suggest_float("drop", 0.15, 0.5) - batch_start = trial.suggest_float("batch_start", 2.0, 8.0) - batch_end = trial.suggest_float("batch_end", 16.0, 64.0) - batch_compound = trial.suggest_float("batch_compound", 1.001, 1.01, log=True) - learn_rate = trial.suggest_float("learn_rate", 1e-4, 5e-3, log=True) - hidden_width = trial.suggest_categorical("hidden_width", [32, 64, 128]) - ner_depth = trial.suggest_int("ner_depth", 2, 6) - ner_width = trial.suggest_categorical("ner_width", [64, 96, 128]) - ner_maxout = trial.suggest_categorical("ner_maxout", [2, 3]) - ner_embed_size = trial.suggest_categorical("ner_embed_size", [2000, 5000, 10000]) - ner_window_size = trial.suggest_int("ner_window_size", 1, 2) - - nlp = _build_pipeline_with_params( - hidden_width=hidden_width, - ner_depth=ner_depth, - ner_width=ner_width, - ner_maxout=ner_maxout, - ner_embed_size=ner_embed_size, - ner_window_size=ner_window_size, - ) - - train_examples = prepare_examples(nlp, train_data) - - best_f1 = _train_for_trial( - nlp, train_examples, dev_data, - n_epochs=tuning_epochs, - drop=drop, - batch_start=batch_start, - batch_end=batch_end, - batch_compound=batch_compound, - learn_rate=learn_rate, - patience=4, - trial=trial, - ) - return best_f1 - - pruner = optuna.pruners.MedianPruner(n_startup_trials=3, n_warmup_steps=3) - study = optuna.create_study( - study_name=study_name, - storage=storage, - direction="maximize", - pruner=pruner, - load_if_exists=True, - ) - study.optimize(objective, n_trials=n_trials, show_progress_bar=True) - - logger.info("=" * 60) - logger.info("BEST TRIAL: #%d F1=%.4f", study.best_trial.number, study.best_value) - for k, v in study.best_params.items(): - logger.info(" %s = %s", k, v) - logger.info("=" * 60) - - return study.best_params - - -def train_with_best_params( - data_path: str, - best_params: Dict, - dev_ratio: float = 0.2, - n_epochs: int = 30, - output_dir: str = "pii_spacy_model_v5", - full_train: bool = False, - patience: int = 7, -): - """Train the final model v5 using the best hyperparameters from Optuna.""" - data = load_train_data(data_path) - - nlp = _build_pipeline_with_params( - hidden_width=best_params.get("hidden_width", 64), - ner_depth=best_params.get("ner_depth", 4), - ner_width=best_params.get("ner_width", 96), - ner_maxout=best_params.get("ner_maxout", 3), - ner_embed_size=best_params.get("ner_embed_size", 2000), - ner_window_size=best_params.get("ner_window_size", 1), - ) - - if full_train: - small_dev = data[:200] - nlp = train_pipeline( - nlp, data, small_dev, - n_epochs=n_epochs, - batch_size_start=best_params.get("batch_start", 4.0), - batch_size_end=best_params.get("batch_end", 32.0), - drop=best_params.get("drop", 0.35), - patience=patience, - output_dir=output_dir, - ) - else: - train_data, dev_data = split_data(data, dev_ratio) - nlp = train_pipeline( - nlp, train_data, dev_data, - n_epochs=n_epochs, - batch_size_start=best_params.get("batch_start", 4.0), - batch_size_end=best_params.get("batch_end", 32.0), - drop=best_params.get("drop", 0.35), - patience=patience, - output_dir=output_dir, - ) - logger.info("=== Final evaluation (v5) on dev set ===") - nlp_best = spacy.load(output_dir) - scores = evaluate_on_data(nlp_best, dev_data) - print_per_label_metrics(scores) - - return nlp - - # ===================================================================== # EVALUATION (strict micro-averaged F1) # ===================================================================== @@ -1114,46 +896,8 @@ def main(): parser.add_argument("--dev-ratio", type=float, default=0.2, help="Dev split ratio") parser.add_argument("--full-train", action="store_true", help="Train on full dataset (no dev split) for final submission") - parser.add_argument("--tune", action="store_true", - help="Run Optuna hyperparameter tuning") - parser.add_argument("--tune-and-train", action="store_true", - help="Tune hyperparameters, then train final model v5") - parser.add_argument("--n-trials", type=int, default=20, - help="Number of Optuna trials") - parser.add_argument("--tuning-epochs", type=int, default=15, - help="Epochs per trial during tuning") - parser.add_argument("--best-params", type=str, default=None, - help="Path to JSON with best params (skip tuning)") args = parser.parse_args() - if args.tune or args.tune_and_train: - if args.best_params: - with open(args.best_params, "r") as f: - best_params = json.load(f) - logger.info("Loaded best params from %s", args.best_params) - else: - best_params = run_hyperparameter_tuning( - data_path=args.data, - dev_ratio=args.dev_ratio, - n_trials=args.n_trials, - tuning_epochs=args.tuning_epochs, - ) - params_path = "best_params_v5.json" - with open(params_path, "w") as f: - json.dump(best_params, f, indent=2, ensure_ascii=False) - logger.info("Best params saved to %s", params_path) - - if args.tune_and_train: - train_with_best_params( - data_path=args.data, - best_params=best_params, - dev_ratio=args.dev_ratio, - n_epochs=args.epochs, - output_dir="pii_spacy_model_v5", - full_train=args.full_train, - ) - return - if args.regex_only: logger.info("=== Regex-only evaluation ===") data = load_train_data(args.data)