NERGAL / nergal.py
ppuzio's picture
Fix public tokenizer batch shape before 1.0.2 release
510b295
Raw
History Blame Contribute Delete
12.1 kB
"""NERGAL hybrid PII cleaner: frozen regex ∪ windowed XLM-R BIO head.
This file is the public PII island. It does not import the lab training stack and
must not call embedding-extension. The packed tokenizer already has the gap ids.
"""
from __future__ import annotations
import hashlib
import json
import math
import re
from dataclasses import dataclass
from functools import lru_cache
from pathlib import Path
import scrub_pii
from scrub_pii import PHONE_TAG, PII_TAG
HUB_ID = 'SlayerLab/NERGAL'
VERSION = '1.0.2'
GAPS = ['[PII_SPACE]', '[PII_BREAK]']
GAP_IDS = [250002, 250003]
BIO_LABELS = ['O', 'B-phone', 'I-phone', 'B-pii', 'I-pii']
LABELS = ['phone', 'pii']
THRESHOLD = 0.95
RULES_SHA = '3016ae5bd403ff997458f9dd74bad8c6ed1388eb83dadc1b31cdb182f9ed607f'
def sha(path):
with Path(path).open('rb') as stream:
return hashlib.file_digest(stream, 'sha256').hexdigest()
def verify_rules(path=None):
digest = sha(path or scrub_pii.__file__)
if digest != RULES_SHA:
raise ValueError(f'Unexpected rules sha256 {digest}')
return digest
@dataclass(frozen=True)
class Unit:
model: str
start: int
end: int
gap: bool
def unitize(text, encode=None, unk=None):
units = []
for match in re.finditer(r'\s+|\S', text):
raw = match.group()
gap = raw.isspace()
model = GAPS[int(any(c in raw for c in '\r\n\v\f\x85\u2028\u2029'))] if gap else raw
if not gap and encode is not None and not encode(model):
if not unk:
raise ValueError('Zero-piece unit without unknown token')
model = unk
units.append(Unit(model, match.start(), match.end(), gap))
return units
def windows(units, count, *, max_units=384, limit=512):
overlap = 128
result, start = [], 0
while start < len(units):
lo, hi = start + 1, min(start + max_units, len(units))
while lo < hi:
mid = (lo + hi + 1) // 2
if count(units[start:mid]) <= limit:
lo = mid
else:
hi = mid - 1
end, size = lo, count(units[start:lo])
if size > limit or (end < len(units) and end - start <= overlap):
raise ValueError('Token budget cannot fit a progressing window')
width = 64
result.append({
'start': start, 'end': end, 'tokens': size,
'owner_start': start if not result else start + width - 1,
'owner_end': end if end == len(units) else end - width + 1,
})
if end == len(units):
break
start = end - overlap
return result
def raw_span(units, a, b, label, score):
if not 0 <= a < b <= len(units) or label not in LABELS or not math.isfinite(score) or not 0 <= score <= 1:
raise ValueError('Invalid unit prediction')
while a < b and units[a].gap:
a += 1
while a < b and units[b - 1].gap:
b -= 1
if a == b:
return None
return {'start': units[a].start, 'end': units[b - 1].end, 'label': label, 'score': score}
def decode_bio(units, logits):
if len(units) != len(logits) or any(len(v) != 5 or any(not math.isfinite(x) for x in v) for v in logits):
raise ValueError('Invalid BIO logits')
result, active, probabilities = [], None, []
def finish(end):
nonlocal active
if active is None:
return
start, label = active
span = raw_span(units, start, end, label, min(probabilities[start:end]))
if span is not None:
result.append(span)
active = None
for i, values in enumerate(logits):
tag = max(range(5), key=lambda j: values[j])
exponentials = [math.exp(x - max(values)) for x in values]
probabilities.append(exponentials[tag] / sum(exponentials))
label = LABELS[(tag - 1) // 2] if tag else None
if tag == 0 or tag in (1, 3) or active is None or active[1] != label:
finish(i)
active = (i, label) if tag else None
finish(len(units))
return result
def decode(spans, threshold=THRESHOLD):
if any(not math.isfinite(s['score']) or not 0 <= s['score'] <= 1 for s in spans):
raise ValueError('Nonfinite/invalid confidence')
result = []
for span in sorted(spans, key=lambda s: (-s['score'], -(s['end'] - s['start']), s['start'], s['label'])):
if span['score'] >= threshold and not any(span['start'] < p['end'] and p['start'] < span['end'] for p in result):
result.append(span)
return sorted(result, key=lambda s: (s['start'], s['end'], s['label']))
class Encoding:
def __init__(self, tokenizer):
self.tokenizer = tokenizer
tokenizer.model_max_length = 512
self.pieces = lru_cache(maxsize=16384)(lambda s: tuple(tokenizer.encode(s, add_special_tokens=False)))
def encode(self, words):
encoded = self.tokenizer([words], is_split_into_words=True, truncation=False, padding=False)
ids = encoded['input_ids'][0]
mapping = encoded.word_ids(0)
first, actual = {}, {}
for i, word in enumerate(mapping):
if word is not None:
first.setdefault(word, i)
actual.setdefault(word, []).append(ids[i])
if set(first) != set(range(len(words))):
raise ValueError('Tokenizer dropped a unit')
if any(tuple(actual[j]) != self.pieces(word) for j, word in enumerate(words)):
raise ValueError('Unit token IDs change with window context')
return encoded, [first[j] for j in range(len(words))]
def count(self, units):
encoded, _ = self.encode([u.model for u in units])
return len(encoded['input_ids'][0])
def prepare(self, text):
units = unitize(text, self.pieces, self.tokenizer.unk_token)
return units, windows(units, self.count) if units else []
def rules(text):
verify_rules()
result = []
scrub_pii.scrub_pii(text, spans=result)
if '[PII]' in text or '[Telefon]' in text:
return []
return sorted(({k: s[k] for k in ('start', 'end', 'label')} | {'score': 1.0} for s in result),
key=lambda s: s['start'])
def apply_union(text, spans):
labels = [None] * len(text)
for span in spans:
start, end, label = span['start'], span['end'], span['label']
if not 0 <= start < end <= len(text):
raise ValueError('Span outside text')
for i in range(start, end):
if labels[i] is None or label == 'phone':
labels[i] = label
out, chars, n_phone, n_pii, i = [], 0, 0, 0, 0
while i < len(text):
lab = labels[i]
if lab is None:
out.append(text[i])
i += 1
continue
j = i + 1
while j < len(text) and labels[j] == lab:
j += 1
tag = PHONE_TAG if lab == 'phone' else PII_TAG
out.append(tag)
chars += len(tag)
if lab == 'phone':
n_phone += 1
else:
n_pii += 1
i = j
return ''.join(out), chars, n_phone, n_pii
def scrub_spans(text, rule_spans, model_spans, *, threshold=THRESHOLD):
rule_keys = {(s['start'], s['end'], s['label']) for s in rule_spans}
model_keep = decode(model_spans, threshold)
extra = sum(1 for s in model_keep if (s['start'], s['end'], s['label']) not in rule_keys)
_, rules_chars, _, _ = apply_union(text, rule_spans)
masked, union_chars, n_phone, n_pii = apply_union(text, list(rule_spans) + model_keep)
return masked, {
'phone': n_phone,
'pii': n_pii,
'rules_placeholder_chars': rules_chars,
'union_placeholder_chars': union_chars,
'model_extra_spans': extra,
}
def _resolve(source, *, local_files_only):
path = Path(source)
if path.is_dir():
return path
from huggingface_hub import snapshot_download
return Path(snapshot_download(source, local_files_only=local_files_only))
def _load_rules_module(asset):
path = Path(asset) / 'scrub_pii.py'
if path.is_file():
import importlib.util
spec = importlib.util.spec_from_file_location('_nergal_scrub_pii', path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
verify_rules(path)
return module
verify_rules()
return scrub_pii
class Nergal:
def __init__(self, asset, device='cpu'):
import torch
from transformers import AutoModelForTokenClassification, AutoTokenizer
self.device = device
self._torch = torch
asset = Path(asset)
self._scrub = _load_rules_module(asset)
card = json.loads((asset / 'hybrid.json').read_text())
if card['gap_ids'] != GAP_IDS or card['threshold'] != THRESHOLD:
raise ValueError('hybrid.json does not match this NERGAL snapshot')
tokenizer = AutoTokenizer.from_pretrained(
str(asset), local_files_only=True, use_fast=True, fix_mistral_regex=False,
)
if [tokenizer.convert_tokens_to_ids(t) for t in GAPS] != GAP_IDS:
raise ValueError('Packed NERGAL tokenizer is missing gap ids')
self.model = AutoModelForTokenClassification.from_pretrained(str(asset), local_files_only=True)
self.encoding = Encoding(tokenizer)
self.threshold = THRESHOLD
self.model.to(device).eval()
@classmethod
def from_pretrained(cls, source=HUB_ID, *, device=None, local_files_only=False):
import torch
if device is None:
device = 'mps' if torch.backends.mps.is_available() else 'cpu'
return cls(_resolve(source, local_files_only=local_files_only), device=device)
def predict(self, text):
torch = self._torch
units, chunks = self.encoding.prepare(text)
if not units:
return []
sums, counts = torch.zeros(len(units), 5), torch.zeros(len(units), 1)
with torch.inference_mode():
for window in chunks:
a, b = window['start'], window['end']
words = [u.model for u in units[a:b]]
encoded, first = self.encoding.encode(words)
batch = self.encoding.tokenizer.pad(
[{k: v[0] for k, v in encoded.items()}], padding=True, return_tensors='pt',
)
batch = {k: v.to(self.device) if torch.is_tensor(v) else v for k, v in batch.items()}
if batch['input_ids'].shape[1] > 512:
raise ValueError('Batch exceeds encoder limit')
logits = self.model(**batch).logits
sums[a:b] += logits[0, first].float().cpu()
counts[a:b] += 1
if (counts == 0).any():
raise ValueError('Missing inference units')
return decode_bio(units, (sums / counts).tolist())
def rule_spans(self, text):
result = []
self._scrub.scrub_pii(text, spans=result)
if '[PII]' in text or '[Telefon]' in text:
return []
return sorted(({k: s[k] for k in ('start', 'end', 'label')} | {'score': 1.0} for s in result),
key=lambda s: s['start'])
def scrub(self, text):
if not text:
return text, {'phone': 0, 'pii': 0, 'rules_placeholder_chars': 0,
'union_placeholder_chars': 0, 'model_extra_spans': 0}
return scrub_spans(text, self.rule_spans(text), self.predict(text), threshold=self.threshold)
def main(argv=None):
import argparse
import sys
parser = argparse.ArgumentParser(description='NERGAL hybrid PII cleaner')
parser.add_argument('--repo', default=HUB_ID)
parser.add_argument('--device', default=None)
parser.add_argument('--local', action='store_true')
args = parser.parse_args(argv)
nergal = Nergal.from_pretrained(args.repo, device=args.device, local_files_only=args.local)
text = sys.stdin.read()
masked, counts = nergal.scrub(text)
sys.stdout.write(masked)
print(json.dumps(counts), file=sys.stderr)
if __name__ == '__main__':
main()