ppuzio commited on
Commit
97bea30
·
verified ·
1 Parent(s): caa10a8

Add hybrid PII scripts: nergal.py and frozen scrub_pii.py

Browse files
Files changed (1) hide show
  1. nergal.py +326 -0
nergal.py ADDED
@@ -0,0 +1,326 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """NERGAL hybrid PII cleaner: frozen regex ∪ windowed XLM-R BIO head.
2
+
3
+ This file is the public PII island. It does not import the lab training stack and
4
+ must not call embedding-extension. The packed tokenizer already has the gap ids.
5
+ """
6
+ from __future__ import annotations
7
+
8
+ import hashlib
9
+ import json
10
+ import math
11
+ import re
12
+ from dataclasses import dataclass
13
+ from functools import lru_cache
14
+ from pathlib import Path
15
+
16
+ import scrub_pii
17
+ from scrub_pii import PHONE_TAG, PII_TAG
18
+
19
+ HUB_ID = 'SlayerLab/NERGAL'
20
+ GAPS = ['[PII_SPACE]', '[PII_BREAK]']
21
+ GAP_IDS = [250002, 250003]
22
+ BIO_LABELS = ['O', 'B-phone', 'I-phone', 'B-pii', 'I-pii']
23
+ LABELS = ['phone', 'pii']
24
+ THRESHOLD = 0.95
25
+ RULES_SHA = '547c0428b0799bf051566d6ac489987eff27f09e1bda36a5665452fe155b3966'
26
+
27
+
28
+ def sha(path):
29
+ with Path(path).open('rb') as stream:
30
+ return hashlib.file_digest(stream, 'sha256').hexdigest()
31
+
32
+
33
+ def verify_rules(path=None):
34
+ digest = sha(path or scrub_pii.__file__)
35
+ if digest != RULES_SHA:
36
+ raise ValueError(f'Unexpected rules sha256 {digest}')
37
+ return digest
38
+
39
+
40
+ @dataclass(frozen=True)
41
+ class Unit:
42
+ model: str
43
+ start: int
44
+ end: int
45
+ gap: bool
46
+
47
+
48
+ def unitize(text, encode=None, unk=None):
49
+ units = []
50
+ for match in re.finditer(r'\s+|\S', text):
51
+ raw = match.group()
52
+ gap = raw.isspace()
53
+ model = GAPS[int(any(c in raw for c in '\r\n\v\f\x85\u2028\u2029'))] if gap else raw
54
+ if not gap and encode is not None and not encode(model):
55
+ if not unk:
56
+ raise ValueError('Zero-piece unit without unknown token')
57
+ model = unk
58
+ units.append(Unit(model, match.start(), match.end(), gap))
59
+ return units
60
+
61
+
62
+ def windows(units, count, *, max_units=384, limit=512):
63
+ overlap = 128
64
+ result, start = [], 0
65
+ while start < len(units):
66
+ lo, hi = start + 1, min(start + max_units, len(units))
67
+ while lo < hi:
68
+ mid = (lo + hi + 1) // 2
69
+ if count(units[start:mid]) <= limit:
70
+ lo = mid
71
+ else:
72
+ hi = mid - 1
73
+ end, size = lo, count(units[start:lo])
74
+ if size > limit or (end < len(units) and end - start <= overlap):
75
+ raise ValueError('Token budget cannot fit a progressing window')
76
+ width = 64
77
+ result.append({
78
+ 'start': start, 'end': end, 'tokens': size,
79
+ 'owner_start': start if not result else start + width - 1,
80
+ 'owner_end': end if end == len(units) else end - width + 1,
81
+ })
82
+ if end == len(units):
83
+ break
84
+ start = end - overlap
85
+ return result
86
+
87
+
88
+ def raw_span(units, a, b, label, score):
89
+ if not 0 <= a < b <= len(units) or label not in LABELS or not math.isfinite(score) or not 0 <= score <= 1:
90
+ raise ValueError('Invalid unit prediction')
91
+ while a < b and units[a].gap:
92
+ a += 1
93
+ while a < b and units[b - 1].gap:
94
+ b -= 1
95
+ if a == b:
96
+ return None
97
+ return {'start': units[a].start, 'end': units[b - 1].end, 'label': label, 'score': score}
98
+
99
+
100
+ def decode_bio(units, logits):
101
+ if len(units) != len(logits) or any(len(v) != 5 or any(not math.isfinite(x) for x in v) for v in logits):
102
+ raise ValueError('Invalid BIO logits')
103
+ result, active, probabilities = [], None, []
104
+
105
+ def finish(end):
106
+ nonlocal active
107
+ if active is None:
108
+ return
109
+ start, label = active
110
+ span = raw_span(units, start, end, label, min(probabilities[start:end]))
111
+ if span is not None:
112
+ result.append(span)
113
+ active = None
114
+
115
+ for i, values in enumerate(logits):
116
+ tag = max(range(5), key=lambda j: values[j])
117
+ exponentials = [math.exp(x - max(values)) for x in values]
118
+ probabilities.append(exponentials[tag] / sum(exponentials))
119
+ label = LABELS[(tag - 1) // 2] if tag else None
120
+ if tag == 0 or tag in (1, 3) or active is None or active[1] != label:
121
+ finish(i)
122
+ active = (i, label) if tag else None
123
+ finish(len(units))
124
+ return result
125
+
126
+
127
+ def decode(spans, threshold=THRESHOLD):
128
+ if any(not math.isfinite(s['score']) or not 0 <= s['score'] <= 1 for s in spans):
129
+ raise ValueError('Nonfinite/invalid confidence')
130
+ result = []
131
+ for span in sorted(spans, key=lambda s: (-s['score'], -(s['end'] - s['start']), s['start'], s['label'])):
132
+ if span['score'] >= threshold and not any(span['start'] < p['end'] and p['start'] < span['end'] for p in result):
133
+ result.append(span)
134
+ return sorted(result, key=lambda s: (s['start'], s['end'], s['label']))
135
+
136
+
137
+ class Encoding:
138
+ def __init__(self, tokenizer):
139
+ self.tokenizer = tokenizer
140
+ tokenizer.model_max_length = 512
141
+ self.pieces = lru_cache(maxsize=16384)(lambda s: tuple(tokenizer.encode(s, add_special_tokens=False)))
142
+
143
+ def encode(self, words):
144
+ encoded = self.tokenizer(words, is_split_into_words=True, truncation=False, padding=False)
145
+ ids = encoded['input_ids'][0]
146
+ mapping = encoded.word_ids(0)
147
+ first, actual = {}, {}
148
+ for i, word in enumerate(mapping):
149
+ if word is not None:
150
+ first.setdefault(word, i)
151
+ actual.setdefault(word, []).append(ids[i])
152
+ if set(first) != set(range(len(words))):
153
+ raise ValueError('Tokenizer dropped a unit')
154
+ if any(tuple(actual[j]) != self.pieces(word) for j, word in enumerate(words)):
155
+ raise ValueError('Unit token IDs change with window context')
156
+ return encoded, [first[j] for j in range(len(words))]
157
+
158
+ def count(self, units):
159
+ encoded, _ = self.encode([u.model for u in units])
160
+ return len(encoded['input_ids'][0])
161
+
162
+ def prepare(self, text):
163
+ units = unitize(text, self.pieces, self.tokenizer.unk_token)
164
+ return units, windows(units, self.count) if units else []
165
+
166
+
167
+ def rules(text):
168
+ verify_rules()
169
+ result = []
170
+ scrub_pii.scrub_pii(text, spans=result)
171
+ if '[PII]' in text or '[Telefon]' in text:
172
+ return []
173
+ return sorted(({k: s[k] for k in ('start', 'end', 'label')} | {'score': 1.0} for s in result),
174
+ key=lambda s: s['start'])
175
+
176
+
177
+ def apply_union(text, spans):
178
+ labels = [None] * len(text)
179
+ for span in spans:
180
+ start, end, label = span['start'], span['end'], span['label']
181
+ if not 0 <= start < end <= len(text):
182
+ raise ValueError('Span outside text')
183
+ for i in range(start, end):
184
+ if labels[i] is None or label == 'phone':
185
+ labels[i] = label
186
+ out, chars, n_phone, n_pii, i = [], 0, 0, 0, 0
187
+ while i < len(text):
188
+ lab = labels[i]
189
+ if lab is None:
190
+ out.append(text[i])
191
+ i += 1
192
+ continue
193
+ j = i + 1
194
+ while j < len(text) and labels[j] == lab:
195
+ j += 1
196
+ tag = PHONE_TAG if lab == 'phone' else PII_TAG
197
+ out.append(tag)
198
+ chars += len(tag)
199
+ if lab == 'phone':
200
+ n_phone += 1
201
+ else:
202
+ n_pii += 1
203
+ i = j
204
+ return ''.join(out), chars, n_phone, n_pii
205
+
206
+
207
+ def scrub_spans(text, rule_spans, model_spans, *, threshold=THRESHOLD):
208
+ rule_keys = {(s['start'], s['end'], s['label']) for s in rule_spans}
209
+ model_keep = decode(model_spans, threshold)
210
+ extra = sum(1 for s in model_keep if (s['start'], s['end'], s['label']) not in rule_keys)
211
+ _, rules_chars, _, _ = apply_union(text, rule_spans)
212
+ masked, union_chars, n_phone, n_pii = apply_union(text, list(rule_spans) + model_keep)
213
+ return masked, {
214
+ 'phone': n_phone,
215
+ 'pii': n_pii,
216
+ 'rules_placeholder_chars': rules_chars,
217
+ 'union_placeholder_chars': union_chars,
218
+ 'model_extra_spans': extra,
219
+ }
220
+
221
+
222
+ def _resolve(source, *, local_files_only):
223
+ path = Path(source)
224
+ if path.is_dir():
225
+ return path
226
+ from huggingface_hub import snapshot_download
227
+ return Path(snapshot_download(source, local_files_only=local_files_only))
228
+
229
+
230
+ def _load_rules_module(asset):
231
+ path = Path(asset) / 'scrub_pii.py'
232
+ if path.is_file():
233
+ import importlib.util
234
+ spec = importlib.util.spec_from_file_location('_nergal_scrub_pii', path)
235
+ module = importlib.util.module_from_spec(spec)
236
+ spec.loader.exec_module(module)
237
+ verify_rules(path)
238
+ return module
239
+ verify_rules()
240
+ return scrub_pii
241
+
242
+
243
+ class Nergal:
244
+ def __init__(self, asset, device='cpu'):
245
+ import torch
246
+ from transformers import AutoModelForTokenClassification, AutoTokenizer
247
+ self.device = device
248
+ self._torch = torch
249
+ asset = Path(asset)
250
+ self._scrub = _load_rules_module(asset)
251
+ card = json.loads((asset / 'hybrid.json').read_text())
252
+ if card['gap_ids'] != GAP_IDS or card['threshold'] != THRESHOLD:
253
+ raise ValueError('hybrid.json does not match this NERGAL snapshot')
254
+ tokenizer = AutoTokenizer.from_pretrained(
255
+ str(asset), local_files_only=True, use_fast=True, fix_mistral_regex=False,
256
+ )
257
+ if [tokenizer.convert_tokens_to_ids(t) for t in GAPS] != GAP_IDS:
258
+ raise ValueError('Packed NERGAL tokenizer is missing gap ids')
259
+ self.model = AutoModelForTokenClassification.from_pretrained(str(asset), local_files_only=True)
260
+ self.encoding = Encoding(tokenizer)
261
+ self.threshold = THRESHOLD
262
+ self.model.to(device).eval()
263
+
264
+ @classmethod
265
+ def from_pretrained(cls, source=HUB_ID, *, device=None, local_files_only=False):
266
+ import torch
267
+ if device is None:
268
+ device = 'mps' if torch.backends.mps.is_available() else 'cpu'
269
+ return cls(_resolve(source, local_files_only=local_files_only), device=device)
270
+
271
+ def predict(self, text):
272
+ torch = self._torch
273
+ units, chunks = self.encoding.prepare(text)
274
+ if not units:
275
+ return []
276
+ sums, counts = torch.zeros(len(units), 5), torch.zeros(len(units), 1)
277
+ with torch.inference_mode():
278
+ for window in chunks:
279
+ a, b = window['start'], window['end']
280
+ words = [u.model for u in units[a:b]]
281
+ encoded, first = self.encoding.encode(words)
282
+ batch = self.encoding.tokenizer.pad(
283
+ [{k: v[0] for k, v in encoded.items()}], padding=True, return_tensors='pt',
284
+ )
285
+ batch = {k: v.to(self.device) if torch.is_tensor(v) else v for k, v in batch.items()}
286
+ if batch['input_ids'].shape[1] > 512:
287
+ raise ValueError('Batch exceeds encoder limit')
288
+ logits = self.model(**batch).logits
289
+ sums[a:b] += logits[0, first].float().cpu()
290
+ counts[a:b] += 1
291
+ if (counts == 0).any():
292
+ raise ValueError('Missing inference units')
293
+ return decode_bio(units, (sums / counts).tolist())
294
+
295
+ def rule_spans(self, text):
296
+ result = []
297
+ self._scrub.scrub_pii(text, spans=result)
298
+ if '[PII]' in text or '[Telefon]' in text:
299
+ return []
300
+ return sorted(({k: s[k] for k in ('start', 'end', 'label')} | {'score': 1.0} for s in result),
301
+ key=lambda s: s['start'])
302
+
303
+ def scrub(self, text):
304
+ if not text:
305
+ return text, {'phone': 0, 'pii': 0, 'rules_placeholder_chars': 0,
306
+ 'union_placeholder_chars': 0, 'model_extra_spans': 0}
307
+ return scrub_spans(text, self.rule_spans(text), self.predict(text), threshold=self.threshold)
308
+
309
+
310
+ def main(argv=None):
311
+ import argparse
312
+ import sys
313
+ parser = argparse.ArgumentParser(description='NERGAL hybrid PII cleaner')
314
+ parser.add_argument('--repo', default=HUB_ID)
315
+ parser.add_argument('--device', default=None)
316
+ parser.add_argument('--local', action='store_true')
317
+ args = parser.parse_args(argv)
318
+ nergal = Nergal.from_pretrained(args.repo, device=args.device, local_files_only=args.local)
319
+ text = sys.stdin.read()
320
+ masked, counts = nergal.scrub(text)
321
+ sys.stdout.write(masked)
322
+ print(json.dumps(counts), file=sys.stderr)
323
+
324
+
325
+ if __name__ == '__main__':
326
+ main()