File size: 7,856 Bytes
4039e77
 
 
 
 
 
 
bbd1d02
4039e77
 
 
 
7ce967b
4039e77
 
bbd1d02
7ce967b
e0caa8b
 
bbd1d02
4039e77
 
 
 
 
 
 
510b295
 
 
 
 
 
 
 
 
 
 
78cc85a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2bf5c25
 
 
 
 
 
 
4464d47
 
 
 
 
 
 
 
 
 
 
 
 
e0caa8b
 
 
 
 
 
 
 
 
 
 
 
4039e77
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bbd1d02
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4039e77
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
"""Synthetic NERGAL tests. Invented strings only; no corpus text or real identifiers."""
import hashlib
import json
import unittest
from pathlib import Path

HERE = Path(__file__).resolve().parent
RULES_SHA = 'b238d5b88aa3f3d55a24bb051ec93f9179dfb14b2c650441c8e0acb225d81d59'


class NergalTests(unittest.TestCase):
    def test_card_and_rules_hash(self):
        from nergal import GAP_IDS, GAPS, HUB_ID, RULES_SHA as PINNED, THRESHOLD, VERSION
        card = json.loads((HERE / 'hybrid.json').read_text())
        self.assertEqual(HUB_ID, 'SlayerLab/NERGAL')
        self.assertEqual(VERSION, '1.2.0')
        self.assertEqual(card['version'], VERSION)
        self.assertEqual(card['eval']['union_fp'], 80)
        self.assertEqual(card['eval']['rules_fp'], 24)
        self.assertEqual((card['eval']['whole_entities'], card['eval']['gold_entities']), (303, 315))  # restated gold
        self.assertEqual(GAPS, card['gaps'])
        self.assertEqual(GAP_IDS, card['gap_ids'])
        self.assertEqual(THRESHOLD, card['threshold'])
        self.assertEqual(PINNED, RULES_SHA)
        digest = hashlib.sha256((HERE / 'scrub_pii.py').read_bytes()).hexdigest()
        self.assertEqual(digest, RULES_SHA)

    def test_real_tokenizer_preserves_batch_and_unit_alignment(self):
        from transformers import AutoTokenizer
        from nergal import Encoding
        tokenizer = AutoTokenizer.from_pretrained(str(HERE), local_files_only=True, fix_mistral_regex=False)
        encoding = Encoding(tokenizer)
        words = ['A', '[PII_SPACE]', '1']
        encoded, first = encoding.encode(words)
        self.assertIsInstance(encoded['input_ids'][0], list)
        self.assertEqual(len(first), len(words))
        self.assertEqual([encoded.word_ids(0)[i] for i in first], [0, 1, 2])

    def test_window_token_count_matches_the_encoded_window(self):
        from transformers import AutoTokenizer
        from nergal import Encoding
        tokenizer = AutoTokenizer.from_pretrained(str(HERE), local_files_only=True, fix_mistral_regex=False)
        encoding = Encoding(tokenizer)
        text = ' '.join(f'Zdanie {i}: tel. 22 123 45 67,\nNIP 1234567802.' for i in range(120))
        units, chunks = encoding.prepare(text)
        self.assertGreater(len(chunks), 1)
        for w in chunks:
            encoded, _ = encoding.encode([u.model for u in units[w['start']:w['end']]])
            self.assertEqual(w['tokens'], len(encoded['input_ids'][0]))
            self.assertLessEqual(w['tokens'], 512)

    def test_float16_is_opt_in_and_needs_an_accelerator(self):
        from nergal import Nergal
        with self.assertRaises(ValueError):
            Nergal(HERE, device='cpu', dtype='float16')
        with self.assertRaises(ValueError):
            Nergal(HERE, dtype='bfloat16')

    def test_existing_placeholders_do_not_switch_the_rules_off(self):
        from nergal import rules
        text = 'Kontakt [Telefon], NIP 1234567802.'  # invented, checksum-valid
        [span] = rules(text)
        self.assertEqual(text[span['start']:span['end']], '1234567802')
        self.assertEqual(rules('a [PII] b [Telefon] c'), [])

    def test_grouped_national_phones_mask_without_a_cue(self):
        from nergal import rules
        for text, masked in (('Sklep Ala, 601 234 567, czynne 9-17', ['601 234 567']),
                             ('Biuro: (22) 123 45 67.', ['(22) 123 45 67']),
                             ('Zapraszamy: +48 601 234 567.', ['+48 601 234 567']),
                             ('Zapraszamy: 601234567.', []),               # plain 9 digits stay cue-gated
                             ('Budżet wyniósł 601 234 567 zł.', []),       # amount
                             ('Wartość 601 234 567,89 w tabeli.', []),     # decimal figure
                             ('Kwota 500 000 000 osób.', []),              # round count
                             ('NIP: 601 234 567', [])):                    # other-number label
            with self.subTest(text=text):
                self.assertEqual([text[s['start']:s['end']] for s in rules(text) if s['label'] == 'phone'], masked)

    def test_email_ends_at_glued_text_and_mention_lists_are_not_addresses(self):
        from nergal import rules
        for text, masked in (('kontakt@firma.plKontakt', ['kontakt@firma.pl']),     # capital glued onto the TLD
                             ('jan@firma.plwww.firma.pl', ['jan@firma.pl']),        # URL host glued onto the TLD
                             ('jan@firma.plkontakt', ['jan@firma.plkontakt']),      # all-lowercase glue: known limit
                             ('BIURO@FIRMA.PL', ['BIURO@FIRMA.PL']),
                             ('kontakt@jan7@wp.pl', ['jan7@wp.pl']),                # word glued on with '@'
                             ('Dzięki @kasia @firma.pl @tomek', []),                # list of mentions
                             ('Obserwuj @jan@firma.social', ['jan@firma.social'])):  # handle keeps the mask
            with self.subTest(text=text):
                self.assertEqual([text[s['start']:s['end']] for s in rules(text)], masked)

    def test_union_keeps_regex_and_adds_model_spans(self):
        from nergal import apply_union, scrub_spans
        text = 'Ring 000000000 then extra.'
        rules = [{'start': 5, 'end': 14, 'label': 'phone', 'score': 1.0}]
        model = [
            {'start': 5, 'end': 14, 'label': 'phone', 'score': 0.99},
            {'start': 20, 'end': 25, 'label': 'pii', 'score': 0.97},
        ]
        masked, counts = scrub_spans(text, rules, model, threshold=0.95)
        self.assertIn('[Telefon]', masked)
        self.assertIn('[PII]', masked)
        self.assertGreater(counts['union_placeholder_chars'], counts['rules_placeholder_chars'])
        self.assertEqual(counts['model_extra_spans'], 1)
        _, rules_chars, _, _ = apply_union(text, rules)
        self.assertEqual(counts['rules_placeholder_chars'], rules_chars)
        self.assertNotIn('000000000', masked)
        self.assertNotIn('extra', masked)

    def test_model_phone_spans_follow_the_phone_policy(self):
        from nergal import model_keep, scrub_spans
        span = lambda text, part, label='phone', score=0.99: dict(
            start=text.index(part), end=text.index(part) + len(part), label=label, score=score)
        for text, part, kept in (('tel. 112', '112', []),                                  # emergency number
                                 ('tel. 51 23 45', '51 23 45', []),                        # under 7 digits
                                 ('tel. 601 234 567/602 345 678', '601 234 567/602 345 678',
                                  ['601 234 567', '602 345 678']),                         # one span per number
                                 ('tel. 601 234 567, fax, 602 345 678', '601 234 567, fax, 602 345 678',
                                  ['601 234 567', '602 345 678']),                         # a word between parts
                                 ('tel. 22 123 45 67 wew. 101', '22 123 45 67 wew. 101',
                                  ['22 123 45 67 wew. 101']),                              # extension stays inside
                                 ('Jan Kowalski, 112', 'Jan Kowalski', ['Jan Kowalski'])):  # other labels unchanged
            label = 'pii' if part[0].isalpha() else 'phone'
            with self.subTest(text=text):
                keep = model_keep(text, [span(text, part, label)])
                self.assertEqual([text[s['start']:s['end']] for s in keep], kept)
                self.assertTrue(all(s['label'] == label and s['score'] == 0.99 for s in keep))
        text = 'tel. 112'
        self.assertEqual(scrub_spans(text, [], [span(text, '112')])[0], text)
        self.assertEqual(model_keep(text, [span(text, '112', score=0.9)], threshold=0.95), [])


if __name__ == '__main__':
    unittest.main()