Code-Blue/crater-detection
0
1import os2import cv23import torchvision4import torch5import numpy as np6from torchvision import transforms7import gradio8import requests9 10examples_path = [11 ["crtr.jpg"],12 ["crtr1.jpg"],13 ["crtr2.png"],14 ["crttr3.jpg"],15 ["crtr4.webp"],16]17 18# examples = []19 20# def load_examples():21# for ex in examples_path:22# ex_img = cv2.imread(ex)23# ex_img = cv2.cvtColor(ex_img,cv2.COLOR_BGR2RGB)24# examples.append([ex_img])25 26 27model = torch.load("FasterRCNN_model PT.pth",weights_only = False, map_location=torch.device('cpu'))28device = torch.device('cpu')29 30def show_preds_image(image_path):31 img = cv2.imread(image_path)32 img = cv2.cvtColor(img,cv2.COLOR_BGR2RGB)33 out_val = model([transforms.ToTensor()(img).to(device)])34 out_bbox = out_val[0]['boxes']35 out_scores = out_val[0]['scores'] 36 keep = torchvision.ops.nms(out_bbox,out_scores,0.45).cpu().detach().numpy()37 out_bbox = out_bbox.cpu().detach().numpy()38 39 for box in out_bbox[keep]:40 x1,y1,x2,y2 = int(box[0]), int(box[1]), int(box[2]), int(box[3])41 cv2.rectangle(img,(x1,y1),(x2,y2),(255,255,255),2,lineType=cv2.LINE_AA)42 43 return img44 45inputs_image = [gradio.components.Image(type='filepath',label='Input Image')]46outputs_image = [gradio.components.Image(type='numpy',label="Output Image")]47 48interface_image = gradio.Interface(49 fn=show_preds_image,50 inputs=inputs_image,51 outputs=outputs_image,52 title="Lunar crater detector",53 examples=examples_path,54 cache_examples=False55)56 57 58gradio.TabbedInterface([interface_image],tab_names=['Image Interface']).queue().launch()