NERGAL / nergal.py
ppuzio's picture
Claude Opus 5.5
Prepare 1.1.0: batch API (predict_many, scrub_many), opt-in float16
fdc79fc
Raw History Blame
14.2 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.1.0'
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 = 'f32d5c5452fc47178e109d4bc248a0d8234ea6e59e8cf79407f4eb8451581d67'
DTYPES = ('float32', 'float16')
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):
"""Special tokens + cached unit pieces. Exact because encode() rejects any unit whose pieces change with
context; no re-tokenization inside the windows() binary search."""
return self.tokenizer.num_special_tokens_to_add() + sum(len(self.pieces(u.model)) for u in units)
def prepare(self, text):
units = unitize(text, self.pieces, self.tokenizer.unk_token)
return units, windows(units, self.count) if units else []
def spans_from(module, text):
"""Rule spans on the original text. Existing [PII]/[Telefon] placeholders do not switch the rules off."""
result = []
module.scrub_pii(text, spans=result)
return sorted(({k: s[k] for k in ('start', 'end', 'label')} | {'score': 1.0} for s in result),
key=lambda s: s['start'])
def rules(text):
verify_rules()
return spans_from(scrub_pii, text)
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', dtype='float32'):
import torch
from transformers import AutoModelForTokenClassification, AutoTokenizer
if dtype not in DTYPES:
raise ValueError(f'dtype must be one of {DTYPES}')
if dtype == 'float16' and torch.device(device).type == 'cpu':
raise ValueError('float16 needs a cuda or mps device')
self.dtype = dtype
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, dtype=getattr(torch, dtype)).eval()
@classmethod
def from_pretrained(cls, source=HUB_ID, *, device=None, dtype='float32', local_files_only=False):
import torch
if device is None:
device = 'cuda' if torch.cuda.is_available() else 'mps' if torch.backends.mps.is_available() else 'cpu'
return cls(_resolve(source, local_files_only=local_files_only), device=device, dtype=dtype)
def predict(self, text):
return self.predict_many([text])[0]
def predict_many(self, texts, *, batch_tokens=32768, max_batch=128, group=64):
"""Model spans for each text. Windows of up to `group` texts are sorted by token length and packed into
batches of at most `max_batch` rows and `batch_tokens` padded tokens; per text, units, windows, averaging
and decode are the same as one window at a time."""
if min(batch_tokens, max_batch, group) < 1:
raise ValueError('batch_tokens, max_batch and group must be positive')
texts = list(texts)
result = []
for i in range(0, len(texts), group):
result += self._predict_group(texts[i:i + group], batch_tokens, max_batch)
return result
def _predict_group(self, texts, batch_tokens, max_batch):
torch, encoding = self._torch, self.encoding
docs, items = [], []
for d, text in enumerate(texts):
units, chunks = encoding.prepare(text)
docs.append((units, torch.zeros(len(units), 5), torch.zeros(len(units), 1)))
for window in chunks:
encoded, first = encoding.encode([u.model for u in units[window['start']:window['end']]])
items.append((len(encoded['input_ids'][0]), d, window, encoded, first))
items.sort(key=lambda item: item[0])
with torch.inference_mode():
i = 0
while i < len(items):
j = i + 1 # sorted ascending, so items[j] sets the padded width of items[i:j + 1]
while j < len(items) and j - i < max_batch and (j - i + 1) * items[j][0] <= batch_tokens:
j += 1
part = items[i:j]
batch = encoding.tokenizer.pad([{k: v[0] for k, v in e.items()} for _, _, _, e, _ in part],
padding=True, return_tensors='pt')
if batch['input_ids'].shape[1] > 512:
raise ValueError('Batch exceeds encoder limit')
logits = self.model(**{k: v.to(self.device) for k, v in batch.items()}).logits.float().cpu()
for row, (_, d, window, _, first) in enumerate(part):
a, b = window['start'], window['end']
docs[d][1][a:b] += logits[row, first]
docs[d][2][a:b] += 1
i = j
result = []
for units, sums, counts in docs:
if (counts == 0).any():
raise ValueError('Missing inference units')
result.append(decode_bio(units, (sums / counts).tolist()) if units else [])
return result
def rule_spans(self, text):
return spans_from(self._scrub, text)
def scrub(self, text):
return self.scrub_many([text])[0]
def scrub_many(self, texts):
texts = list(texts)
return [scrub_spans(text, self.rule_spans(text), model, threshold=self.threshold)
for text, model in zip(texts, self.predict_many(texts), strict=True)]
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('--dtype', default='float32', choices=DTYPES)
parser.add_argument('--local', action='store_true')
args = parser.parse_args(argv)
nergal = Nergal.from_pretrained(args.repo, device=args.device, dtype=args.dtype,
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()