import spaces # MUST be first — ZeroGPU patches torch.cuda before transformers imports torch import torch import gradio as gr import json from transformers import pipeline as hf_pipeline MODEL_ID = "genzeonplatform/cliniguard-diagnosis-icd-ner" # Entity type -> human-readable description (9 diagnosis/coding categories) ENTITY_INFO = { "PRIMARY_DIAGNOSIS": "Principal reason for encounter", "SECONDARY_DIAGNOSIS": "Additional diagnoses", "DIFFERENTIAL_DIAGNOSIS": "Diagnoses under consideration", "COMORBIDITY": "Co-existing conditions", "COMPLICATION": "Hospital-acquired or treatment complications", "CHRONIC_CONDITION": "Long-term conditions", "ACUTE_CONDITION": "Acute episodes or events", "DIAGNOSIS_STATUS": "Current clinical status (active, resolved, improving)", "DIAGNOSIS_DATE": "Date of diagnosis or onset", } # Color map for HighlightedText (9 distinct colors) COLOR_MAP = { "PRIMARY_DIAGNOSIS": "#ef4444", "SECONDARY_DIAGNOSIS": "#f97316", "DIFFERENTIAL_DIAGNOSIS": "#eab308", "COMORBIDITY": "#22c55e", "COMPLICATION": "#3b82f6", "CHRONIC_CONDITION": "#a855f7", "ACUTE_CONDITION": "#ec4899", "DIAGNOSIS_STATUS": "#14b8a6", "DIAGNOSIS_DATE": "#6b7280", } # Load the token-classification pipeline at module scope. # aggregation_strategy="simple" merges B-/I- subword tokens into full entity spans, # exactly as documented in the model card. nlp = hf_pipeline( "token-classification", model=MODEL_ID, aggregation_strategy="simple", ) # Move the underlying model to CUDA eagerly (ZeroGPU hijack intercepts this). nlp.model.to("cuda") @spaces.GPU(duration=60) def extract_entities(text: str): """Extract diagnosis entities from unstructured clinical text. Recognizes 9 diagnosis/coding entity types: PRIMARY_DIAGNOSIS, SECONDARY_DIAGNOSIS, DIFFERENTIAL_DIAGNOSIS, COMORBIDITY, COMPLICATION, CHRONIC_CONDITION, ACUTE_CONDITION, DIAGNOSIS_STATUS, and DIAGNOSIS_DATE. Args: text: Clinical text to analyze (discharge summaries, progress notes, etc.) Returns: A tuple of (highlighted_text, structured_json, entity_summary_table). """ # Run inference; pipeline handles tokenization + aggregation raw_entities = nlp(text) # Build HighlightedText format: {"text": ..., "entities": [{"start","end","entity"}]} entities = [ { "start": ent["start"], "end": ent["end"], "entity": ent["entity_group"], } for ent in raw_entities ] highlighted = {"text": text, "entities": entities} # Build structured JSON output structured = [ { "text": ent["word"], "type": ent["entity_group"], "description": ENTITY_INFO.get(ent["entity_group"], ""), "score": round(float(ent["score"]), 4), "start": ent["start"], "end": ent["end"], } for ent in raw_entities ] # Build entity summary table (type, count, examples, description) summary = {} for ent in raw_entities: label = ent["entity_group"] if label not in summary: summary[label] = [] summary[label].append(ent["word"]) table_data = [] for label, mentions in sorted(summary.items()): table_data.append([ label, len(mentions), ", ".join(mentions[:5]) + ("..." if len(mentions) > 5 else ""), ENTITY_INFO.get(label, ""), ]) return highlighted, json.dumps(structured, indent=2), table_data # --- Gradio UI --- CSS = """ #col-container { max-width: 1100px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ EXAMPLES = [ ["Discharge Diagnosis: Primary: acute myocardial infarction. " "Secondary: type 2 diabetes mellitus, essential hypertension. " "Status: improving. Complication: acute kidney injury."], ["Patient admitted with community-acquired pneumonia. " "PMH: COPD, coronary artery disease, hyperlipidemia. " "Differential: pulmonary embolism vs pneumonia. Status: suspected. Date: 03/15/2024."], ["Assessment: sepsis is worsening with acute respiratory failure. " "Chronic conditions include type 2 diabetes and chronic kidney disease stage 3. " "Comorbidity: obesity. Complication: hospital-acquired pneumonia."], ["Principal Dx: congestive heart failure with acute respiratory failure. " "Secondary: atrial fibrillation, chronic kidney disease. " "Status: active. Date of onset: 01/20/2024."], ] with gr.Blocks(theme=gr.themes.Citrus(), css=CSS) as demo: gr.Markdown( """ # 🏥 CliniGuard Diagnosis ICD NER Extract diagnosis entities — primary/secondary/differential diagnoses, comorbidities, complications, chronic/acute conditions, status, and dates — from unstructured clinical text using a **PubMedBERT**-based token-classification model ([genzeonplatform/cliniguard-diagnosis-icd-ner](https://huggingface.co/genzeonplatform/cliniguard-diagnosis-icd-ner)). """ ) with gr.Column(elem_id="col-container"): with gr.Row(): input_text = gr.Textbox( label="Clinical Text", placeholder="Enter clinical text (discharge summaries, progress notes, clinical narratives)...", lines=8, scale=4, ) run_btn = gr.Button("Extract Entities", variant="primary", scale=1) highlighted_output = gr.HighlightedText( label="Annotated Clinical Text", combine_adjacent=True, show_legend=True, color_map=COLOR_MAP, ) with gr.Accordion("Structured Output", open=False): json_output = gr.JSON(label="Extracted Entities (JSON)") entity_table = gr.Dataframe( headers=["Entity Type", "Count", "Examples", "Description"], label="Entity Summary", interactive=False, wrap=True, ) gr.Examples( examples=EXAMPLES, inputs=[input_text], outputs=[highlighted_output, json_output, entity_table], fn=extract_entities, cache_examples=True, cache_mode="lazy", ) run_btn.click( fn=extract_entities, inputs=[input_text], outputs=[highlighted_output, json_output, entity_table], api_name="extract_entities", ) if __name__ == "__main__": demo.launch(mcp_server=True)