NERGAL / test_nergal.py
ppuzio's picture
Claude Opus 5.5
2.0.1: a [PERSON] placeholder ends an other-number label
960c515
Raw History Blame Contribute Delete
15 kB
"""Synthetic NERGAL tests. Invented strings only; no corpus text or real identifiers."""
import hashlib
import json
import shutil
import tempfile
import unittest
from pathlib import Path
HERE = Path(__file__).resolve().parent
RULES_SHA = 'c1b924a893ed01b739d6616fcc1e05ffc5c9c0df138ac1409ab3af427c6b0768'
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, '2.0.1')
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
for marker in ('[PHONE]', '[Telefon]', '[PII]', '[PERSON]'):
with self.subTest(marker=marker):
text = f'Kontakt {marker}, NIP 1234567802.' # invented, checksum-valid
[span] = rules(text)
self.assertEqual(text[span['start']:span['end']], '1234567802')
self.assertEqual(rules('a [PII] b [PHONE] c [PERSON] d [Telefon] e'), [])
def test_phones_are_tagged_phone(self):
from nergal import apply_union, rules
text = 'Biuro: (22) 123 45 67.'
masked = apply_union(text, rules(text))[0]
self.assertEqual(masked, 'Biuro: [PHONE].')
def test_new_and_legacy_phone_tags_are_the_same_boundary(self):
from nergal import rules
def found(text): # values, not offsets: the two tags differ in length
return [(text[s['start']:s['end']], s['label']) for s in rules(text)]
for text in ('Telefon: {} lub 601234567', 'Kontakt: {}, 601234567', # plain 9 digits: cue-gated
'tel. {}\nwew. 123 Jan Nowak\nwew. 456 sekretariat',
'Kontakt {}: e-mail biuro@example.pl, 601 234 567'): # invented
with self.subTest(text=text):
self.assertEqual(found(text.format('[Telefon]')), found(text.format('[PHONE]')))
def test_a_name_placeholder_ends_an_other_number_label_but_never_a_phone_cue(self):
from nergal import rules
for text, masked in (('Kod [PERSON] zadzwoń 601 200 300', ['601 200 300']),
('Kod Jan zadzwoń 601 200 300', []), # a raw name is text to the rules
('[PERSON] NIP: 601 234 567', []), # a label after the name still counts
('Telefon do [PERSON]: 601234567', ['601234567']), # plain 9 digits: cue-gated
('Kontakt: [PERSON], 601234567', ['601234567']),
('Informacje u [PERSON] pod numerem 601234567', ['601234567']),
('tel. [PERSON] 601234567', ['601234567'])): # invented
with self.subTest(text=text):
self.assertEqual([text[s['start']:s['end']] for s in rules(text) if s['label'] == 'phone'], masked)
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('[PHONE]', 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.assertEqual(counts['person'], 0)
self.assertNotIn('000000000', masked)
self.assertNotIn('extra', masked)
def test_union_ranks_phone_over_pii_over_person(self):
from nergal import apply_union
text = 'Jan Kowal 601 234 567'
spans = [{'start': 0, 'end': 21, 'label': 'person'}, {'start': 4, 'end': 9, 'label': 'pii'},
{'start': 10, 'end': 21, 'label': 'phone'}]
masked, chars, counts = apply_union(text, spans)
self.assertEqual(masked, '[PERSON][PII][PERSON][PHONE]') # 0–4 and the space at 9 stay person
self.assertEqual(chars, len(masked))
self.assertEqual(counts, {'person': 2, 'pii': 1, 'phone': 1})
with self.assertRaises(ValueError):
apply_union(text, [{'start': 0, 'end': 1, 'label': 'org'}])
def test_person_spans_expand_merge_and_skip_markers(self):
from nergal import person_spans
text = 'Pani Nowakowskiej-Kowal, O’Brien. [PERSON] i [PHONE]'
a = text.index('Nowak'); b = text.index('O’B')
spans = person_spans(text, [(a + 2, a + 6), (b, b + 2), (text.index('[PERSON]') + 1, text.index('[PERSON]') + 3)])
self.assertEqual([text[s['start']:s['end']] for s in spans], ['Nowakowskiej-Kowal', 'O’Brien'])
self.assertEqual(person_spans('Jan Nowak', [(0, 3), (4, 9)]), [{'start': 0, 'end': 9, 'label': 'person', 'score': 1.0}])
self.assertEqual(person_spans('Jan Nowak', [(0, 3), (4, 9)])[0]['end'], 9) # no-break space joins
self.assertEqual(person_spans('Jan, Nowak', [(0, 3), (5, 10)])[1]['start'], 5) # punctuation keeps them apart
self.assertEqual(len(person_spans('Jan\nNowak', [(0, 3), (4, 9)])), 2) # a line break keeps them apart
self.assertEqual(person_spans('Nowak', [(0, 2), (1, 5)]), [{'start': 0, 'end': 5, 'label': 'person', 'score': 1.0}])
self.assertEqual(person_spans('x', []), [])
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), [])
def _names(self):
from nergal import Names
return Names(HERE, json.loads((HERE / 'hybrid.json').read_text())['names'])
def test_names_mask_an_invented_person_and_nothing_else(self):
names = self._names()
text = 'Wniosek złożyła Anna Nowakowska z Radomia.'
self.assertIn('Anna Nowakowska', [text[s['start']:s['end']] for s in names.spans(text)])
self.assertEqual(names.spans('Zdanie bez nazwisk o pogodzie.'), []) # special tokens are never a person
self.assertEqual(names.spans(''), [])
text = 'Ala ma kota.' # a first name alone is a person (names policy)
self.assertEqual([text[s['start']:s['end']] for s in names.spans(text)], ['Ala'])
def test_names_find_a_person_past_the_first_window_and_are_idempotent(self):
from nergal import apply_union
names = self._names()
text = 'Zdanie bez nazwisk o pogodzie. ' * 200 + 'Wniosek złożyła Anna Nowakowska.'
spans = names.spans(text)
self.assertEqual([text[s['start']:s['end']] for s in spans], ['Anna Nowakowska'])
masked, _, _ = apply_union(text, spans)
self.assertEqual(names.spans(masked), [])
def test_names_batched_equal_one_at_a_time(self):
names = self._names()
texts = ['Wniosek złożyła Anna Nowakowska z Radomia.', '',
'Zdanie bez nazwisk o pogodzie. ' * 200 + 'Podpisał Tomasz Wrzos.', 'Ala ma kota.']
self.assertEqual(names.spans_many(texts, max_batch=2), [names.spans(t) for t in texts])
def test_names_are_opt_in_and_verified(self):
from nergal import Names
card = json.loads((HERE / 'hybrid.json').read_text())['names']
with self.assertRaises(ValueError):
Names(HERE, {**card, 'sha256': {**card['sha256'], 'tokenizer.json': '0' * 64}})
def test_a_snapshot_without_names_refuses_names(self):
from nergal import Nergal
card = json.loads((HERE / 'hybrid.json').read_text())
del card['names']
with tempfile.TemporaryDirectory() as tmp:
shutil.copy(HERE / 'scrub_pii.py', tmp)
(Path(tmp) / 'hybrid.json').write_text(json.dumps(card))
with self.assertRaisesRegex(ValueError, 'no names model'):
Nergal(tmp, names=True) # raised before the tokenizer or model loads
def test_hub_download_skips_names_unless_asked(self):
from unittest import mock
import huggingface_hub
from nergal import _resolve
with mock.patch.object(huggingface_hub, 'snapshot_download', return_value=str(HERE)) as download:
_resolve('SlayerLab/NERGAL', local_files_only=True)
_resolve('SlayerLab/NERGAL', local_files_only=True, names=True)
self.assertEqual([c.kwargs['ignore_patterns'] for c in download.call_args_list], [['names/*'], None])
def test_pack_keeps_every_index_once_within_limits(self):
from nergal import pack
lengths = [5, 1, 9, 3, 3, 7]
parts = pack(lengths, batch_tokens=12, max_batch=2)
self.assertEqual(sorted(i for p in parts for i in p), list(range(6)))
for p in parts:
self.assertLessEqual(len(p), 2)
self.assertLessEqual(len(p) * max(lengths[i] for i in p), 12)
if __name__ == '__main__':
unittest.main()