-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdeidentify_debug.py
More file actions
124 lines (102 loc) · 2.97 KB
/
Copy pathdeidentify_debug.py
File metadata and controls
124 lines (102 loc) · 2.97 KB
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
#!/usr/bin/env python3
from transformers import pipeline
import torch
import re
import argparse
from pathlib import Path
# -----------------------
# Arguments
# -----------------------
ap = argparse.ArgumentParser()
ap.add_argument(
"--input",
required=True,
help="Path to input transcription text file (.txt)"
)
ap.add_argument(
"--output",
default=None,
help="Optional output file path. If not set, '_output_deidentified.txt' is used."
)
ap.add_argument(
"--chunk_size",
type=int,
default=300,
help="Chunk size for BERT (default: 300)"
)
args = ap.parse_args()
input_file = Path(args.input).expanduser().resolve()
print("DEBUG: using input_file =", input_file)
if not input_file.exists():
raise FileNotFoundError(f"Input file not found: {input_file}")
if args.output:
output_file = Path(args.output).expanduser().resolve()
else:
output_file = input_file.with_name(
input_file.stem + "_output_deidentified.txt"
)
CHUNK_SIZE = args.chunk_size
# -----------------------
# Load NER model
# -----------------------
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
ner = pipeline(
"ner",
model="lm2445/for_deidentify",
grouped_entities=True,
device=0 if torch.cuda.is_available() else -1
)
# -----------------------
# Read input text
# -----------------------
text = input_file.read_text(encoding="utf-8").strip()
# -----------------------
# Chunking (BERT 512 limit)
# -----------------------
def chunk_text(text, size):
chunks = []
start = 0
while start < len(text):
end = min(start + size, len(text))
chunks.append(text[start:end])
start = end
return chunks
chunks = chunk_text(text, CHUNK_SIZE)
# -----------------------
# Process chunks with NER + replacement
# -----------------------
processed_chunks = []
for chunk in chunks:
results = ner(chunk)
deidentified = chunk
for entity in results:
if "word" not in entity:
continue
ent_text = entity["word"]
if "entity_group" in entity and entity["entity_group"]:
ent_label = entity["entity_group"]
elif "entity" in entity and entity["entity"]:
ent_label = entity["entity"].split("-")[-1]
else:
continue
placeholder = f"<{ent_label}>"
# Escape entity text to avoid regex issues
pattern = re.escape(ent_text)
deidentified = re.sub(pattern, placeholder, deidentified)
processed_chunks.append(deidentified)
# -----------------------
# Combine chunks
# -----------------------
deidentified_text = "".join(processed_chunks)
# -----------------------
# Save output
# -----------------------
output_file.write_text(deidentified_text, encoding="utf-8")
print(f"De-identified text saved to: {output_file}")
# -----------------------
# GPU memory
# -----------------------
if torch.cuda.is_available():
peak = torch.cuda.max_memory_allocated()
print(f"Peak GPU memory usage: {peak / 1024**2:.2f} MB")