Spanbertapi / app.py
2001haitem's picture
Update app.py
160ad2c verified
Raw History Blame Contribute Delete
2.83 kB
import gradio as gr
from transformers import pipeline, AutoTokenizer
# ==============================
# πŸ”§ Load model and tokenizer
# ==============================
MODEL_NAME = "hazarri/fine_tuned_spanbert"
ner_pipeline = pipeline(
"token-classification",
model=MODEL_NAME,
aggregation_strategy="simple"
)
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
# ==============================
# 🧠 1️⃣ ADR Extraction Function
# ==============================
def extract_adrs(text):
results = ner_pipeline(text)
if not results:
return []
return [r["word"] for r in results]
# ==============================
# βœ‚οΈ 2️⃣ Tokenizer Function
# ==============================
def tokenize_text(text):
tokens = tokenizer.tokenize(text)
return tokens
# ==============================
# 🎨 Gradio Interfaces
# ==============================
adr_interface = gr.Interface(
fn=extract_adrs,
inputs=gr.Textbox(lines=4, placeholder="Enter a medical review..."),
outputs=gr.JSON(label="Extracted ADRs"),
title="πŸ’Š ADR Extraction (SpanBERT)",
description="Extracts adverse drug reactions (ADRs) from text using a fine-tuned SpanBERT model.",
api_name="/predict" # βœ… for remote calls via gradio_client
)
tokenizer_interface = gr.Interface(
fn=tokenize_text,
inputs=gr.Textbox(lines=3, placeholder="Enter text to tokenize..."),
outputs=gr.JSON(label="Tokens"),
title="πŸ”€ Tokenizer",
description="Displays tokens produced by the model's tokenizer.",
api_name="/tokenize" # βœ… second endpoint
)
# Combine both tools
demo = gr.TabbedInterface(
[adr_interface, tokenizer_interface],
["ADR Extraction", "Tokenizer"]
)
# ==============================
# πŸš€ Launch the app
# ==============================
if __name__ == "__main__":
demo.launch()
# import gradio as gr
# from transformers import pipeline
# # Load your fine-tuned SpanBERT ADR model
# ner_pipeline = pipeline(
# "token-classification",
# model="hazarri/fine_tuned_spanbert", # replace with your actual model name
# aggregation_strategy="simple"
# )
# def extract_adrs(text):
# # Run the model
# results = ner_pipeline(text)
# # Return empty list if no entities detected
# if not results:
# return []
# # Extract only the ADR words
# adrs = [r["word"] for r in results]
# return adrs
# # Gradio interface
# demo = gr.Interface(
# fn=extract_adrs,
# inputs=gr.Textbox(lines=5, placeholder="Enter a medical review..."),
# outputs=gr.JSON(label="Extracted ADRs"),
# title="πŸ’Š SpanBERT ADR Extraction API",
# description="Extracts adverse drug reactions (ADRs) from patient reviews using a fine-tuned SpanBERT model."
# )
# if __name__ == "__main__":
# demo.launch()