import torch from PIL import Image import gradio as gr from open_clip import create_model_from_pretrained, get_tokenizer MODEL_NAME = "microsoft/BiomedCLIP-PubMedBERT_256-vit_base_patch16_224" print("Loading BiomedCLIP...") model, preprocess = create_model_from_pretrained(f"hf-hub:{MODEL_NAME}") tokenizer = get_tokenizer(f"hf-hub:{MODEL_NAME}") model.eval() print("Model Loaded") BASE_PROMPTS = { "Eczema": "a dermatology clinical image showing eczema with dry itchy inflamed patches", "Psoriasis": "a dermatology image of psoriasis with thick scaly plaques", "Fungal Infection (Tinea)": "a fungal skin infection with circular expanding ring like rash", "Acne": "a dermatology image showing acne with pimples and inflammation", "Dermatitis": "a dermatitis rash with redness irritation and inflammation", "Urticaria (Hives)": "raised itchy welts on skin like urticaria or hives", "Benign Mole": "a harmless benign mole on the skin", "Melanoma Suspicion": "a suspicious melanoma skin lesion with asymmetry border irregularity dark color", "Healthy Skin": "a healthy normal skin image with no lesions" } CONFIDENCE_THRESHOLD = 0.42 # tune if needed def build_prompts(symptoms): prompts = [] for label, base in BASE_PROMPTS.items(): if symptoms and symptoms.strip() != "": enriched_prompt = f"{base}. Patient symptoms: {symptoms}" else: enriched_prompt = base prompts.append((label, enriched_prompt)) return prompts def predict(image, symptoms): if image is None: return "Please upload an image.", None image = preprocess(image).unsqueeze(0) prompts = build_prompts(symptoms) labels = [p[0] for p in prompts] text_list = [p[1] for p in prompts] with torch.no_grad(): image_features = model.encode_image(image) text_tokens = tokenizer(text_list) text_features = model.encode_text(text_tokens) image_features /= image_features.norm(dim=-1, keepdim=True) text_features /= text_features.norm(dim=-1, keepdim=True) similarity = (100.0 * image_features @ text_features.T).softmax(dim=-1) probs = similarity.squeeze().tolist() label_scores = list(zip(labels, probs)) label_scores.sort(key=lambda x: x[1], reverse=True) best_label, best_prob = label_scores[0] explanation = "" if best_prob < CONFIDENCE_THRESHOLD: explanation = ( "⚠️ The model is uncertain about this case. " "Consider consulting a dermatologist, especially if symptoms are worsening, painful, rapidly spreading, " "bleeding, or changing in shape/color." ) best_label = "Uncertain — Needs Clinical Evaluation" top3 = {label: round(score, 3) for label, score in label_scores[:3]} return ( f"Prediction: {best_label} (confidence: {round(best_prob,3)})\n\n" f"{explanation}\n\n" "🔒 Disclaimer: This is a research tool and NOT a medical diagnosis.", top3 ) ui = gr.Interface( fn=predict, inputs=[ gr.Image(type="pil", label="Upload Skin Image"), gr.Textbox(label="Describe Symptoms (optional)", placeholder="e.g., itchy red rash for 2 weeks, burning sensation, spreading, no bleeding") ], outputs=[ gr.Textbox(label="Result"), gr.Label(num_top_classes=3, label="Top Predictions") ], title="BiomedCLIP Dermatology Assistant", description="Upload a skin image and optionally describe symptoms. Uses zero-shot BiomedCLIP for prediction. Research use only — not medical advice." ) ui.launch()