gimmy256's picture
Update app.py
2d639e3 verified
Raw
History Blame Contribute Delete
7.95 kB
"""
MedASR + Gemini LLM Pipeline — HuggingFace Space
=================================================
Required secret: GEMINI_API_KEY
"""
import os, re, json
import gradio as gr
import google.generativeai as genai
import onnxruntime as ort
import numpy as np
import soundfile as sf
# Safe API configuration
api_key = os.environ.get("GEMINI_API_KEY")
if api_key:
genai.configure(api_key=api_key)
SAMPLES = {
"GP Consultation — Fever & headache (Eka)":
"i have been having fever and headache for the past three days and also some body ache "
"antigen test for dengue was done and it came negative so i am prescribing you dolo 650 "
"for fever and pantop dsr 40 for your stomach please take them after food",
"GP Consultation — Thyroid symptoms (Eka)":
"patient feels sad not able to wake up and energy less blood test was done thyroid levels "
"came low so starting thyroxine 25 mcg once daily in the morning on empty stomach "
"follow up after 15 days",
"Radiology Report — CT Chest PE (MedASR official)":
"ct chest pe protocol indication 54 year old female shortness of breath evaluate for pe "
"technique standard protocol findings pulmonary vasculature the main pa is patent there "
"are filling defects in the segmental branches of the right lower lobe compatible with "
"acute pe no saddle embolus lungs no pneumothorax small bilateral effusions right greater "
"than left impression acute segmental pe right lower lobe",
"Patient Symptoms — Back & leg pain (Hani89)":
"i am experiencing severe back pain and leg pain since last week i was given dolo 650 "
"tablet thrice a day for 6 days and pantop dsr for 7 days and also thyroxine 25 mcg "
"for 15 days",
}
MODEL_PATH = "model_quantized.onnx" # adjust if different
session = None
def load_model():
global session
if session is None:
session = ort.InferenceSession(MODEL_PATH)
return session
def preprocess_audio(audio_path):
audio, sr = sf.read(audio_path)
if len(audio.shape) > 1:
audio = np.mean(audio, axis=1)
# Resample if needed
if sr != 16000:
from scipy.signal import resample
audio = resample(audio, int(len(audio) * 16000 / sr))
return audio.astype(np.float32)
def transcribe_audio(audio_path):
try:
session = load_model()
audio = preprocess_audio(audio_path)
inputs = {session.get_inputs()[0].name: audio}
outputs = session.run(None, inputs)
# This part depends on your model decoding
transcript = str(outputs[0])
return transcript
except Exception as e:
return f"ASR Error: {e}"
def run_full_pipeline(audio):
if not audio:
return "No audio", "", "", ""
transcript = transcribe_audio(audio)
if transcript.startswith("ASR Error"):
return transcript, "", "", ""
correction = run_correction(transcript)
soap = run_soap(transcript)
entities = run_entities(transcript)
return transcript, correction, soap, entities
def prompt_correction(t):
return f"""You are a medical transcript editor.
Fix transcription errors, add punctuation, expand abbreviations, and preserve all clinical content exactly.
ASR transcript:
{t}
"""
def prompt_soap(t):
return f"""Convert this into a structured SOAP note.
SUBJECTIVE:
OBJECTIVE:
ASSESSMENT:
PLAN:
Transcript:
{t}
"""
def prompt_entities(t):
return (
"Extract medical entities and return ONLY valid JSON:\n"
'{"diagnoses":[],"medications":[],"dosages":[],"symptoms":[],"procedures":[],"instructions":[]}\n\n'
f"Transcript:\n{t}"
)
def call_gemini(prompt: str) -> str:
api_key = os.environ.get("GEMINI_API_KEY")
if not api_key:
return "❌ GEMINI_API_KEY not set in Hugging Face Space."
try:
genai.configure(api_key=api_key)
# safer model
model = genai.GenerativeModel("gemini-2.5-flash")
response = model.generate_content(prompt)
return response.text if hasattr(response, "text") else str(response)
except Exception as e:
return f"❌ Gemini API error: {e}"
def load_sample(name):
return SAMPLES.get(name, "")
def run_correction(transcript):
if not transcript.strip():
return "Please enter or select a transcript first."
return call_gemini(prompt_correction(transcript))
def run_soap(transcript):
if not transcript.strip():
return "Please enter or select a transcript first."
return call_gemini(prompt_soap(transcript))
def run_entities(transcript):
if not transcript.strip():
return "Please enter or select a transcript first."
raw = call_gemini(prompt_entities(transcript))
try:
clean = re.sub(r"```json|```", "", raw).strip()
obj = json.loads(clean)
icons = {
"diagnoses": "🔴 Diagnoses",
"medications": "💊 Medications",
"dosages": "📏 Dosages",
"symptoms": "🤒 Symptoms",
"procedures": "🔬 Procedures",
"instructions": "📋 Instructions",
}
lines = []
for key, label in icons.items():
vals = obj.get(key, [])
if vals:
lines.append(f"**{label}**")
for v in vals:
lines.append(f" • {v}")
lines.append("")
return "\n".join(lines) if lines else "No entities found."
except Exception:
return raw
def run_all(transcript):
if not transcript.strip():
empty = "No transcript provided."
return empty, empty, empty
return (
run_correction(transcript),
run_soap(transcript),
run_entities(transcript),
)
with gr.Blocks(
title="MedASR + Gemini Pipeline",
theme=gr.themes.Soft(),
) as demo:
audio_input = gr.Audio(type="filepath", label="Upload Medical Audio")
run_audio_btn = gr.Button("🎤 Run Full Pipeline")
raw_transcript_out = gr.Textbox(label="Raw Transcript", lines=6)
gr.Markdown("""
# ⚕ MedASR + Gemini Pipeline
Phase 2 — Transcript → Clinical NLP (Correction, SOAP, Entities)
""")
with gr.Row():
with gr.Column(scale=1):
gr.Markdown("### Input Transcript")
sample_dd = gr.Dropdown(
choices=list(SAMPLES.keys()),
value=list(SAMPLES.keys())[0],
label="Load sample"
)
transcript_box = gr.Textbox(
value=SAMPLES[list(SAMPLES.keys())[0]],
lines=9,
label="Transcript"
)
sample_dd.change(load_sample, inputs=sample_dd, outputs=transcript_box)
run_all_btn = gr.Button("⚡ Run All")
with gr.Column(scale=1):
gr.Markdown("### Output")
with gr.Tab("Correction"):
correction_out = gr.Textbox(lines=12, interactive=False)
run_a_btn = gr.Button("Run A")
with gr.Tab("SOAP"):
soap_out = gr.Textbox(lines=14, interactive=False)
run_b_btn = gr.Button("Run B")
with gr.Tab("Entities"):
entities_out = gr.Markdown()
run_c_btn = gr.Button("Run C")
run_a_btn.click(run_correction, inputs=transcript_box, outputs=correction_out)
run_b_btn.click(run_soap, inputs=transcript_box, outputs=soap_out)
run_c_btn.click(run_entities, inputs=transcript_box, outputs=entities_out)
run_all_btn.click(
run_all,
inputs=transcript_box,
outputs=[correction_out, soap_out, entities_out]
)
run_audio_btn.click(
fn=run_full_pipeline,
inputs=audio_input,
outputs=[raw_transcript_out, correction_out, soap_out, entities_out]
)
if __name__ == "__main__":
demo.launch()