#!/usr/bin/env python3 # -*- coding: utf-8 -*- """Safe inference example for U4RASD/TypePredictor.""" import argparse import json import torch from modeling_type_predictor import NeoAraBERTTypePredictor OPEN_MARKER = "[ENT]" CLOSE_MARKER = "[/ENT]" def mark_entity(sentence: str, start_char: int, end_char: int) -> str: entity = sentence[start_char:end_char] return sentence[:start_char] + " [ENT] " + entity + " [/ENT] " + sentence[end_char:] def windowed_marked_text(sentence: str, start_char: int, end_char: int, context_chars): if context_chars is None: return mark_entity(sentence, start_char, end_char) left = max(0, start_char - int(context_chars)) right = min(len(sentence), end_char + int(context_chars)) window = sentence[left:right] return mark_entity(window, start_char - left, end_char - left) def encode_safely(tokenizer, config, sentence: str, start_char: int, end_char: int): if not (0 <= start_char < end_char <= len(sentence)): raise ValueError("Invalid character span.") open_id = tokenizer.convert_tokens_to_ids(OPEN_MARKER) close_id = tokenizer.convert_tokens_to_ids(CLOSE_MARKER) candidates = config.get("context_candidates", [None, 500, 300, 150, 80, 30, 0]) for context in candidates: text = windowed_marked_text(sentence, start_char, end_char, context) batch = tokenizer( text, return_tensors="pt", truncation=True, max_length=int(config.get("max_length", 512)), ) ids = batch["input_ids"][0].tolist() if ids.count(open_id) == 1 and ids.count(close_id) == 1 and ids.index(open_id) < ids.index(close_id): return batch, text, context raise RuntimeError("Both entity markers could not be preserved after entity-centered truncation.") def main(): parser = argparse.ArgumentParser() parser.add_argument("--model", default="U4RASD/TypePredictor") parser.add_argument("--sentence", required=True) parser.add_argument("--start-char", type=int, required=True) parser.add_argument("--end-char", type=int, required=True) args = parser.parse_args() model, tokenizer, config = NeoAraBERTTypePredictor.from_pretrained(args.model) model.eval() batch, marked_text, context = encode_safely( tokenizer, config, args.sentence, args.start_char, args.end_char ) with torch.no_grad(): logits = model(**batch)["logits"] probabilities = torch.softmax(logits, dim=-1)[0] index = int(probabilities.argmax()) print(json.dumps({ "entity": args.sentence[args.start_char:args.end_char], "predicted_type": config["labels"][index], "confidence": float(probabilities[index]), "context_chars_used": context, "marked_text": marked_text, }, ensure_ascii=False, indent=2)) if __name__ == "__main__": main()