Views
No views yet
1import torch
2from efficientnet_pytorch import EfficientNet
3from PIL import Image
4import torchvision.transforms as transforms
5import cv2
6import numpy as np
7from PIL import ImageEnhance
8
9def histogram_equalization(img):
10 img = np.array(img)
11 img_yuv = cv2.cvtColor(img, cv2.COLOR_RGB2YUV)
12 img_yuv[:, :, 0] = cv2.equalizeHist(img_yuv[:, :, 0])
13 img = cv2.cvtColor(img_yuv, cv2.COLOR_YUV2RGB)
14 return Image.fromarray(img)
15
16def median_filter(img, kernel_size=3):
17 img = np.array(img)
18 img = cv2.medianBlur(img, kernel_size)
19 return Image.fromarray(img)
20
21def enhance_image(img):
22 img = ImageEnhance.Color(img).enhance(1.2) # Adjust color balance
23 img = ImageEnhance.Contrast(img).enhance(1.2) # Adjust contrast
24 img = ImageEnhance.Sharpness(img).enhance(1.2) # Adjust sharpness
25 return img
26
27# Load the model
28model = EfficientNet.from_pretrained('efficientnet-b0')
29num_classes = 4
30model._fc = nn.Sequential(
31 nn.Dropout(p=0.5),
32 nn.Linear(model._fc.in_features, num_classes)
33)
34model.load_state_dict(torch.load('path_to_your_model.pth'))
35model.eval()
36
37# Define transforms
38transform = transforms.Compose([
39 transforms.Resize((224, 224)),
40 transforms.Lambda(histogram_equalization),
41 transforms.Lambda(median_filter),
42 transforms.Lambda(enhance_image),
43 transforms.ToTensor(),
44 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
45])
46
47# Load and preprocess the image
48img = Image.open('path_to_your_image.jpg')
49img = transform(img).unsqueeze(0)
50
51# Predict
52with torch.no_grad():
53 output = model(img)
54 _, predicted = torch.max(output, 1)
55 severity = predicted.item()
56 print(f'Predicted Acne Severity: {severity}')