PollenNER / app.py
DanielNRU's picture
Update app.py
a3e69f6 verified
Raw
History Blame Contribute Delete
6.67 kB
import re
import torch
import gradio as gr
from transformers import (
AutoModelForTokenClassification,
AutoModelForSequenceClassification,
AutoTokenizer,
)
ID2LABEL = {
0: "O",
1: "B-TOPONYM", 2: "I-TOPONYM",
3: "B-MEDICINE", 4: "I-MEDICINE",
5: "B-SYMPTOM", 6: "I-SYMPTOM",
7: "B-ALLERGEN", 8: "I-ALLERGEN",
9: "B-BODY_PART", 10: "I-BODY_PART",
}
LABEL2ID = {v: k for k, v in ID2LABEL.items()}
REL_LABELS = ["has_symptom", "has_medicine", "no_relation"]
REL_ID2LABEL = {i: l for i, l in enumerate(REL_LABELS)}
REL_LABEL2ID = {l: i for i, l in enumerate(REL_LABELS)}
def load_models():
ner_model = AutoModelForTokenClassification.from_pretrained(
"DanielNRU/pollen-ner",
num_labels=len(ID2LABEL),
id2label=ID2LABEL,
label2id=LABEL2ID,
)
ner_tokenizer = AutoTokenizer.from_pretrained("DeepPavlov/rubert-base-cased")
re_model = AutoModelForSequenceClassification.from_pretrained(
"DanielNRU/pollen-re",
num_labels=len(REL_LABELS),
id2label=REL_ID2LABEL,
label2id=REL_LABEL2ID,
)
re_tokenizer = AutoTokenizer.from_pretrained("DeepPavlov/rubert-base-cased")
ner_model.eval()
re_model.eval()
return ner_model, ner_tokenizer, re_model, re_tokenizer
ner_model, ner_tokenizer, re_model, re_tokenizer = load_models()
def split_sentences(text):
return [s.strip() for s in re.split(r"[.!?]", text) if s.strip()]
def predict_entities(text):
inputs = ner_tokenizer(text, return_tensors="pt", return_offsets_mapping=True, truncation=True, max_length=512)
offset_mapping = inputs.pop("offset_mapping")[0]
with torch.no_grad():
preds = ner_model(**inputs).logits.argmax(dim=-1)[0]
entities = []
current = None
for pred, (start, end) in zip(preds, offset_mapping):
start, end = int(start), int(end)
if start == 0 and end == 0:
continue
label = ID2LABEL[pred.item()]
if label.startswith("B-"):
if current:
entities.append(current)
current = {"text": text[start:end], "label": label[2:], "start": start, "end": end}
elif label.startswith("I-") and current and label[2:] == current["label"]:
current["text"] += text[start:end]
current["end"] = end
else:
if current:
entities.append(current)
current = None
if current:
entities.append(current)
return entities
def insert_entity_markers(text, ent1, ent2):
if ent1["start"] < ent2["start"]:
first, second = ent1, ent2
else:
first, second = ent2, ent1
t = (text[:second["start"]]
+ f"[{second['label']}]"
+ text[second["start"]:second["end"]]
+ f"[/{second['label']}]"
+ text[second["end"]:])
t = (t[:first["start"]]
+ f"[{first['label']}]"
+ t[first["start"]:first["end"]]
+ f"[/{first['label']}]"
+ t[first["end"]:])
return t
def predict_relations(text, entities):
relations = []
for sent in split_sentences(text):
sent_start = text.find(sent)
if sent_start == -1:
continue
sent_end = sent_start + len(sent)
ents = [e for e in entities if e["start"] >= sent_start and e["end"] <= sent_end]
for ent1 in ents:
for ent2 in ents:
if ent1 is ent2:
continue
if ent1["label"] == "BODY_PART" and ent2["label"] in ["SYMPTOM", "MEDICINE"]:
inputs = re_tokenizer(insert_entity_markers(text, ent1, ent2), return_tensors="pt", truncation=True, max_length=256)
with torch.no_grad():
pred = re_model(**inputs).logits.argmax(-1).item()
rel = REL_ID2LABEL[pred]
if rel != "no_relation":
relations.append({"head": ent1, "tail": ent2, "relation": rel})
return relations
def analyze_text(text):
if not text or not text.strip():
return [], "", "", "", ""
entities = predict_entities(text)
relations = predict_relations(text, entities)
highlights = []
used = set()
for e in entities:
s = (e["start"], e["end"], e["label"])
if s not in used:
highlights.append((e["text"], e["label"]))
used.add(s)
toponyms = [e["text"] for e in entities if e["label"] == "TOPONYM"]
medicines = [e["text"] for e in entities if e["label"] == "MEDICINE"]
allergens = [e["text"] for e in entities if e["label"] == "ALLERGEN"]
symptoms = [f"{r['head']['text']} {r['tail']['text']}" for r in relations if r["relation"] == "has_symptom" and r["tail"]["label"] == "SYMPTOM"]
return highlights, ", ".join(toponyms), ", ".join(medicines), ", ".join(symptoms), ", ".join(allergens)
EXAMPLES = [
["В Московской области у меня началась аллергия на пыльцу березы, потекли глаза, нос, принимаю Зиртек и Назонекс."],
["У ребенка в Новокузнецке чешутся глаза, уши и течет нос, врач прописал Кромогексал, Назонекс в нос."],
["В Санкт-Петербурге началось цветение ольхи, сильная реакция, принимаю Эриус, но глаза все равно слезятся."],
]
with gr.Blocks() as demo:
gr.Markdown("## PollenNER — извлечение сущностей и отношений")
gr.Markdown("Извлечение топонимов, лекарств, аллергенов, частей тела и симптомов из сообщений пользователей Пыльца Club.")
with gr.Row():
with gr.Column():
input_text = gr.Textbox(label="Текст", lines=5, placeholder="Введите текст для анализа...")
run_btn = gr.Button("Анализировать")
gr.Examples(examples=EXAMPLES, inputs=input_text)
with gr.Column():
out_highlight = gr.HighlightedText(label="Сущности")
out_toponyms = gr.Textbox(label="Топонимы")
out_medicines = gr.Textbox(label="Медицинские препараты")
out_symptoms = gr.Textbox(label="Симптомы")
out_allergens = gr.Textbox(label="Аллергены")
run_btn.click(fn=analyze_text, inputs=input_text, outputs=[out_highlight, out_toponyms, out_medicines, out_symptoms, out_allergens])
demo.launch()