"""Standalone inference script for ModernBERT-large Medical Dataset NER model. Chunks long inputs at 6000 chars (matching training) and merges per-chunk entities. """ import json import torch import torch.nn as nn from transformers import AutoModel, AutoTokenizer class ModernBERTTokenClassifier(nn.Module): def __init__(self, model_name, num_labels=3, dropout=0.1, freeze_layers=0, use_gradient_checkpointing=False): super().__init__() self.encoder = AutoModel.from_pretrained(model_name) self.dropout = nn.Dropout(dropout) self.linear = nn.Linear(self.encoder.config.hidden_size, num_labels) self.num_labels = num_labels def forward(self, input_ids, attention_mask, labels=None): outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask) seq_out = self.dropout(outputs.last_hidden_state) logits = self.linear(seq_out) if labels is not None: loss_fn = nn.CrossEntropyLoss(ignore_index=-100) loss = loss_fn(logits.view(-1, self.num_labels), labels.view(-1)) return loss return torch.argmax(logits, dim=-1) def load_model(model_dir, device="cpu"): with open(f"{model_dir}/config.json", "r") as f: config = json.load(f) model = ModernBERTTokenClassifier( config["base_model"], config["num_labels"], config["hyperparameters"].get("dropout", 0.1), config["hyperparameters"].get("freeze_layers", 0), ) model.load_state_dict( torch.load(f"{model_dir}/best_model.pt", map_location=device) ) model.to(device).eval() tokenizer = AutoTokenizer.from_pretrained(config["base_model"]) id2label = {int(v): k for k, v in config["label2id"].items()} return model, tokenizer, id2label, config def _predict_chunk(text, offset_base, model, tokenizer, id2label, device, max_length=8192): enc = tokenizer( text, return_offsets_mapping=True, add_special_tokens=True, truncation=True, max_length=max_length, return_attention_mask=True, return_tensors="pt" ) ids = enc["input_ids"].to(device) attn = enc["attention_mask"].to(device) offsets = enc["offset_mapping"][0].tolist() wids = enc.word_ids(0) with torch.no_grad(): preds = model(input_ids=ids, attention_mask=attn)[0].cpu().tolist() word_tags = {} seen = set() for i, wid in enumerate(wids): if wid is not None and wid not in seen: seen.add(wid) word_tags[wid] = id2label[preds[i]] word_spans = {} for i, wid in enumerate(wids): if wid is not None: s, e = offsets[i] if e == 0: continue word_spans[wid] = (s, max(word_spans.get(wid, (s, 0))[1], e)) entities, curr_s, curr_e = [], None, None for wid in sorted(word_spans.keys()): tag = word_tags.get(wid, "O") s, e = word_spans[wid] if tag.startswith("B-"): if curr_s is not None: entities.append({ "text": text[curr_s:curr_e], "start": curr_s + offset_base, "end": curr_e + offset_base, "label": "Dataset", }) curr_s, curr_e = s, e elif tag.startswith("I-") and curr_s is not None: curr_e = e else: if curr_s is not None: entities.append({ "text": text[curr_s:curr_e], "start": curr_s + offset_base, "end": curr_e + offset_base, "label": "Dataset", }) curr_s, curr_e = None, None if curr_s is not None: entities.append({ "text": text[curr_s:curr_e], "start": curr_s + offset_base, "end": curr_e + offset_base, "label": "Dataset", }) return entities def predict(text, model, tokenizer, id2label, device="cpu", chunk_size=6000, chunk_overlap=500, max_length=8192): """Predict entities on arbitrarily long text by chunking at the character level.""" if len(text) <= chunk_size: return _predict_chunk(text, 0, model, tokenizer, id2label, device, max_length) results = [] seen_spans = set() start = 0 while start < len(text): end = min(start + chunk_size, len(text)) chunk = text[start:end] for ent in _predict_chunk(chunk, start, model, tokenizer, id2label, device, max_length): key = (ent["start"], ent["end"], ent["text"]) if key not in seen_spans: seen_spans.add(key) results.append(ent) if end == len(text): break start = end - chunk_overlap return results if __name__ == "__main__": import sys d = sys.argv[1] if len(sys.argv) > 1 else "." model, tok, id2l, cfg = load_model(d) print("Loaded ModernBERT model from {}".format(d)) print(" Model: {}".format(cfg["base_model"])) print(" Max seq length: {}".format(cfg.get("max_seq_length", 8192))) print(" Labels: {}".format(cfg.get("bio_labels", list(id2l.values())))) txt = "We evaluated our method on the MIMIC-III dataset." print("\nSample text: {}".format(txt)) ents = predict(txt, model, tok, id2l) print("Entities found: {}".format(len(ents))) for ent in ents: print(" [{}:{}] {} ({})".format( ent["start"], ent["end"], ent["text"], ent["label"] ))