skinalyze / app.py
Hrigved's picture
Update app.py
5428ea5 verified
Raw
History Blame Contribute Delete
1.96 kB
import torch
import torch.nn as nn
from torchvision import models, transforms
from PIL import Image
import gradio as gr
import os
# Load class names
with open("classes.txt", "r") as f:
class_names = [line.strip() for line in f.readlines()]
# Load model architecture
def load_model():
model = models.densenet121(pretrained=True)
from torchvision.models import DenseNet121_Weights
model = models.densenet121(weights=DenseNet121_Weights.IMAGENET1K_V1)
num_features = model.classifier.in_features
model.classifier = nn.Sequential(
nn.Linear(num_features, 512),
nn.ReLU(),
nn.Dropout(0.4),
nn.Linear(512, 256),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(256, len(class_names))
)
model.load_state_dict(torch.load("model_epoch_25.pth", map_location=torch.device("cpu")))
model.eval()
return model
model = load_model()
# Image preprocessing
transform = transforms.Compose([
transforms.Resize((512, 512)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406],
[0.229, 0.224, 0.225])
])
def predict(image):
image = image.convert("RGB")
img_tensor = transform(image).unsqueeze(0) # Add batch dimension
with torch.no_grad():
outputs = model(img_tensor)
probs = torch.nn.functional.softmax(outputs[0], dim=0)
top_prob, top_idx = torch.max(probs, dim=0)
return {class_names[i]: float(probs[i]) for i in range(len(class_names))}
# Gradio UI
title = "Skin Disease Diagnosis"
description = "Upload a skin image and get the most likely diagnosis. Model: DenseNet121 trained on DermNet."
gr.Interface(
fn=predict,
inputs=gr.Image(type="pil"),
outputs=gr.Label(num_top_classes=5),
title=title,
description=description,
examples=[
["example_images/acne.jpg"],
["example_images/eczema.jpg"]
] if os.path.exists("example_images") else None
).launch()