exposureguard-synthrewrite-t5 / inference_synthrewrite.py
vkatg's picture
Update inference_synthrewrite.py
7bb3d92 verified
Raw
History Blame Contribute Delete
2.84 kB
import re
import torch
from transformers import T5ForConditionalGeneration, T5Tokenizer
_TAGS = ["DOB", "MRN", "PHONE", "ADDR"]
_RE_SPACE = re.compile(r'\[(' + '|'.join(_TAGS) + r')-V\s+(\d+)\]')
_RE_NOCLOSE = re.compile(r'\[(' + '|'.join(_TAGS) + r')-V(\d+)[^\]]{0,3}?(?=\s|$|[.,;])')
def _normalize(text: str) -> str:
text = _RE_SPACE.sub(lambda m: f"[{m.group(1)}-V{m.group(2)}]", text)
text = _RE_NOCLOSE.sub(lambda m: f"[{m.group(1)}-V{m.group(2)}]", text)
return text
def load(model_dir: str = "."):
tok = T5Tokenizer.from_pretrained(model_dir)
model = T5ForConditionalGeneration.from_pretrained(model_dir)
model.eval()
return model, tok
def rewrite(
note: str,
risk: float,
modality: str,
version: int,
trigger_reason: str,
model: T5ForConditionalGeneration,
tokenizer: T5Tokenizer,
max_new_tokens: int = 200,
) -> str:
prompt = (
f"rewrite: [risk={risk:.2f}] [modality={modality}] "
f"[version={version}] [trigger={trigger_reason}] "
f"note: {note}"
)
ids = tokenizer(prompt, return_tensors="pt", max_length=256, truncation=True)
with torch.no_grad():
out = model.generate(
ids["input_ids"],
attention_mask=ids["attention_mask"],
max_new_tokens=max_new_tokens,
num_beams=4,
no_repeat_ngram_size=3,
early_stopping=True,
)
# skip_special_tokens=False preserves bracket placeholders
# which are registered as special tokens in this model
raw = tokenizer.decode(out[0], skip_special_tokens=False)
raw = raw.replace("</s>", "").replace("<pad>", "").strip()
return _normalize(raw)
if __name__ == "__main__":
model, tok = load(".")
examples = [
{
"note": "James Smith is a 67-year-old patient presenting with chest pain. "
"DOB 03/22/1955. Contact: 555-0142. MRN: MRN482910. "
"Attending: Dr. Chen at Memorial General.",
"risk": 0.72, "modality": "asr", "version": 1,
"trigger_reason": "cross_modal_linkage",
},
{
"note": "Patient Maria Garcia (DOB 1978-07-14) presents for follow-up.",
"risk": 0.45, "modality": "text", "version": 0,
"trigger_reason": "none",
},
{
"note": "Lab results for Daniel Harris, DOB 05/15/1960. MRN MRN374821. "
"Age 63. Troponin 0.08. Reviewed by Dr. Patel at Cedar Hills Hospital.",
"risk": 0.83, "modality": "waveform_proxy", "version": 2,
"trigger_reason": "exposure_accumulation",
},
]
for ex in examples:
result = rewrite(**ex, model=model, tokenizer=tok)
print(f"Input: {ex['note']}")
print(f"Output: {result}")
print()