| 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, |
| ) |
| |
| |
| 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() |