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