Spaces:
Sleeping
Sleeping
| import torch | |
| import torch.nn as nn | |
| from torchvision import transforms, models | |
| from torchvision.transforms import InterpolationMode | |
| from torch.utils.data import Dataset, DataLoader | |
| import gradio as gr | |
| from src.utils import utils | |
| from PIL import Image | |
| import numpy as np | |
| def main(window_size: int = 224): | |
| # Set the device | |
| np.random.seed(0) | |
| device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") | |
| print("Running on {}".format(device)) | |
| def predict(inp, stride_size, threshold): | |
| # Apply sliding window on the image | |
| windows, pad_size = utils.sliding_window( | |
| inp, (window_size, window_size), (stride_size, stride_size) | |
| ) | |
| # Define the transformations to be applied to the images | |
| transform = transforms.Compose( | |
| [ | |
| transforms.Resize([256], interpolation=InterpolationMode.BICUBIC), | |
| transforms.CenterCrop([224]), | |
| transforms.ToTensor(), # Converts the image to [0.0, 1.0] range | |
| transforms.Normalize( | |
| mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225] | |
| ), | |
| ] | |
| ) | |
| # Custom dataset class | |
| class CustomDataset(Dataset): | |
| def __init__(self, data, transform=None): | |
| self.data = data | |
| self.transform = transform | |
| def __len__(self): | |
| return len(self.data) | |
| def __getitem__(self, idx): | |
| image, position = self.data[idx] | |
| if self.transform: | |
| image = self.transform(image) | |
| return image, position | |
| # Create custom datasets with transformations | |
| dataset = CustomDataset(windows, transform=transform) | |
| dataloader = DataLoader(dataset, batch_size=128, shuffle=False) | |
| # Load a pretrained ResNet model | |
| model = models.resnet50().to(device) | |
| # Modify the output layer directly to match binary classification | |
| model.fc = nn.Sequential( | |
| nn.Linear(model.fc.in_features, 128), | |
| nn.ReLU(inplace=True), | |
| nn.Linear(128, 1), | |
| ).to(device) | |
| # load model weights (Ensure map_location points to the correct device) | |
| model.load_state_dict( | |
| torch.load("./weights/resnet.pth", map_location=device, weights_only=True) | |
| ) | |
| model = model.to(device) | |
| model.eval() | |
| heatmap = [] | |
| for data in dataloader: | |
| images, positions = data | |
| images = images.to(device) | |
| with torch.no_grad(): | |
| outputs = model(images) | |
| for i in range(len(outputs)): | |
| probability = torch.sigmoid(outputs[i]).item() | |
| img = Image.new( | |
| "RGB", | |
| (window_size, window_size), | |
| color=(int(255 * probability), 0, 0), | |
| ) | |
| heatmap.append((img, (positions[0][i], positions[1][i]))) | |
| # overlay the heatmap on the original image | |
| result, _ = utils.reconstruct_image( | |
| heatmap, pad_size, (window_size, window_size) | |
| ) | |
| prob_threshold = round(threshold * 255, 0) | |
| result = result.point( | |
| lambda p: p * 0 if p < prob_threshold else p | |
| ) # set any pixel value less than | |
| # threshold to 0 | |
| reconstructed_image, _ = utils.reconstruct_image( | |
| windows, pad_size, (window_size, window_size) | |
| ) | |
| result = Image.blend(reconstructed_image, result, alpha=0.5) | |
| # slice the image to remove the padding | |
| # result = result.crop((0, 0, inp.width, inp.height)) | |
| return result | |
| # Gradio Interface with stride size input slider | |
| demo = gr.Interface( | |
| fn=predict, | |
| inputs=[ | |
| gr.Image(type="pil"), | |
| gr.Slider(minimum=32, maximum=224, step=32, label="Stride Size"), | |
| gr.Slider(minimum=0.5, maximum=1.0, step=0.1, label="Threshold"), | |
| ], | |
| outputs=gr.Image(type="pil"), | |
| ) | |
| # launch the interface in a new tab | |
| demo.launch(inbrowser=True, share=True, debug=True) | |
| if __name__ == "__main__": | |
| main() | |